Add multimodal module: receipt images to Java records on the real OpenAI, Anthropic and Ollama models against an OCR-backed local server, validation, repair retry, accuracy by photo condition
Co-Authored-By: Claude Sonnet 5.5 <[email protected]> Claude-Session: https://claude.ai/code/session_01JXVi2GMQ7bR5EmbUFdDj7N
This commit is contained in:
@@ -0,0 +1,39 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
/**
|
||||
* An ESTIMATE of how many input tokens an image costs on Anthropic's API, from the rule in their
|
||||
* vision documentation: an image is split into 28 x 28 pixel patches and costs
|
||||
* ceil(width / 28) * ceil(height / 28) tokens; an image larger than the model tier's long-edge
|
||||
* limit or token limit is scaled down first, keeping its aspect ratio. The scaling here is a simple
|
||||
* search for the largest size that fits, so it can differ by a token or two from what the service
|
||||
* does. This is arithmetic from a documented formula, not a measurement of any bill.
|
||||
*/
|
||||
public final class ImageTokens {
|
||||
|
||||
/** Per-tier limits from the documentation: long edge in pixels, maximum visual tokens. */
|
||||
public record Tier(String name, int maxLongEdge, int maxTokens) {
|
||||
|
||||
public static final Tier STANDARD = new Tier("standard (all but newest)", 1568, 1568);
|
||||
public static final Tier HIGH_RES = new Tier("high-resolution (4.7 and later)", 2576, 4784);
|
||||
}
|
||||
|
||||
private ImageTokens() {
|
||||
}
|
||||
|
||||
public static int patches(int w, int h) {
|
||||
return (int) (Math.ceil(w / 28.0) * Math.ceil(h / 28.0));
|
||||
}
|
||||
|
||||
public static int anthropic(int width, int height, Tier tier) {
|
||||
double scale = Math.min(1.0, tier.maxLongEdge() / (double) Math.max(width, height));
|
||||
while (true) {
|
||||
int w = Math.max(1, (int) Math.floor(width * scale));
|
||||
int h = Math.max(1, (int) Math.floor(height * scale));
|
||||
int t = patches(w, h);
|
||||
if (t <= tier.maxTokens()) {
|
||||
return t;
|
||||
}
|
||||
scale *= 0.995;
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,17 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.time.LocalDate;
|
||||
import java.util.List;
|
||||
|
||||
/**
|
||||
* What we want out of a receipt photo. Money is {@link BigDecimal}, never double: a total that is
|
||||
* 0.1 + 0.2 away from the sum of its lines would make the reconciliation check in
|
||||
* {@link ReceiptValidator} useless.
|
||||
*/
|
||||
public record Receipt(String merchant, LocalDate date, String currency, List<Line> items,
|
||||
BigDecimal subtotal, BigDecimal tax, BigDecimal total) {
|
||||
|
||||
public record Line(String description, int quantity, BigDecimal unitPrice, BigDecimal lineTotal) {
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.core.io.ByteArrayResource;
|
||||
import org.springframework.util.MimeType;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
/**
|
||||
* Sends a receipt image to a chat model and gets a {@link Receipt} back. If the numbers do not add
|
||||
* up it asks once more, quoting what was wrong; a second failure is returned as-is so the caller
|
||||
* can send it to a human.
|
||||
*/
|
||||
public class ReceiptExtractor {
|
||||
|
||||
public static final String SYSTEM = "You read receipts. Extract exactly what is printed. "
|
||||
+ "Never guess a number you cannot read; use null instead.";
|
||||
|
||||
public record Result(Receipt receipt, List<String> problems, int attempts) {
|
||||
|
||||
public boolean ok() {
|
||||
return problems.isEmpty();
|
||||
}
|
||||
}
|
||||
|
||||
private final ChatClient client;
|
||||
|
||||
public ReceiptExtractor(ChatClient client) {
|
||||
this.client = client;
|
||||
}
|
||||
|
||||
public Receipt extractOnce(byte[] image, MimeType type) {
|
||||
return client.prompt().system(SYSTEM)
|
||||
.user(u -> u.text("Extract this receipt.").media(type, new ByteArrayResource(image)))
|
||||
.call().entity(Receipt.class);
|
||||
}
|
||||
|
||||
public Result extract(byte[] image) {
|
||||
Receipt first = extractOnce(image, MimeTypeUtils.IMAGE_PNG);
|
||||
List<String> problems = ReceiptValidator.problems(first);
|
||||
if (problems.isEmpty()) {
|
||||
return new Result(first, problems, 1);
|
||||
}
|
||||
Receipt second = client.prompt().system(SYSTEM)
|
||||
.user(u -> u.text("Extract this receipt. A previous attempt had these problems, so re-read the "
|
||||
+ "image carefully: " + String.join("; ", problems))
|
||||
.media(MimeTypeUtils.IMAGE_PNG, new ByteArrayResource(image)))
|
||||
.call().entity(Receipt.class);
|
||||
return new Result(second, ReceiptValidator.problems(second), 2);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,109 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import java.awt.Color;
|
||||
import java.awt.Font;
|
||||
import java.awt.Graphics2D;
|
||||
import java.awt.RenderingHints;
|
||||
import java.awt.geom.AffineTransform;
|
||||
import java.awt.image.BufferedImage;
|
||||
import java.io.ByteArrayOutputStream;
|
||||
import java.io.IOException;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Random;
|
||||
import javax.imageio.ImageIO;
|
||||
|
||||
/**
|
||||
* Draws a receipt as a PNG so the whole pipeline has a real image to chew on. Deterministic: the
|
||||
* same receipt and the same {@link Look} always produce the same bytes (the noise uses a seeded
|
||||
* {@link Random}).
|
||||
*/
|
||||
public final class ReceiptImages {
|
||||
|
||||
/** How the "photo" was taken. {@code scale} shrinks it, {@code noise} speckles it, {@code degrees} tilts it. */
|
||||
public record Look(String name, double scale, double noise, double degrees) {
|
||||
|
||||
public static final Look CLEAN = new Look("clean", 1.0, 0.0, 0.0);
|
||||
public static final Look SMALL = new Look("small (1/3 size)", 0.33, 0.0, 0.0);
|
||||
public static final Look NOISY = new Look("noisy", 1.0, 0.25, 0.0);
|
||||
public static final Look TILTED = new Look("tilted 4 degrees", 1.0, 0.0, 4.0);
|
||||
}
|
||||
|
||||
private ReceiptImages() {
|
||||
}
|
||||
|
||||
public static List<String> text(Receipt r) {
|
||||
List<String> t = new ArrayList<>();
|
||||
t.add(r.merchant());
|
||||
t.add(r.date().toString());
|
||||
t.add("CURRENCY: " + r.currency());
|
||||
t.add("--------------------------------");
|
||||
for (Receipt.Line l : r.items()) {
|
||||
t.add(String.format("%-14s %d x %s %s", l.description(), l.quantity(), l.unitPrice().toPlainString(),
|
||||
l.lineTotal().toPlainString()));
|
||||
}
|
||||
t.add("--------------------------------");
|
||||
t.add("SUBTOTAL " + r.subtotal().toPlainString());
|
||||
t.add("TAX " + r.tax().toPlainString());
|
||||
t.add("TOTAL " + r.total().toPlainString());
|
||||
return t;
|
||||
}
|
||||
|
||||
public static byte[] png(Receipt r, Look look) {
|
||||
List<String> lines = text(r);
|
||||
int base = 30;
|
||||
int w = 640;
|
||||
int h = 60 + lines.size() * (base + 12);
|
||||
BufferedImage page = new BufferedImage(w, h, BufferedImage.TYPE_INT_RGB);
|
||||
Graphics2D g = page.createGraphics();
|
||||
g.setRenderingHint(RenderingHints.KEY_TEXT_ANTIALIASING, RenderingHints.VALUE_TEXT_ANTIALIAS_ON);
|
||||
g.setColor(Color.WHITE);
|
||||
g.fillRect(0, 0, w, h);
|
||||
g.setColor(Color.BLACK);
|
||||
g.setFont(new Font(Font.MONOSPACED, Font.PLAIN, base));
|
||||
int y = 50;
|
||||
for (String s : lines) {
|
||||
g.drawString(s, 20, y);
|
||||
y += base + 12;
|
||||
}
|
||||
g.dispose();
|
||||
BufferedImage out = page;
|
||||
if (look.degrees() != 0) {
|
||||
BufferedImage rot = new BufferedImage(w, h, BufferedImage.TYPE_INT_RGB);
|
||||
Graphics2D rg = rot.createGraphics();
|
||||
rg.setRenderingHint(RenderingHints.KEY_INTERPOLATION, RenderingHints.VALUE_INTERPOLATION_BILINEAR);
|
||||
rg.setColor(Color.WHITE);
|
||||
rg.fillRect(0, 0, w, h);
|
||||
rg.transform(AffineTransform.getRotateInstance(Math.toRadians(look.degrees()), w / 2.0, h / 2.0));
|
||||
rg.drawImage(page, 0, 0, null);
|
||||
rg.dispose();
|
||||
out = rot;
|
||||
}
|
||||
if (look.scale() != 1.0) {
|
||||
int nw = (int) Math.round(w * look.scale());
|
||||
int nh = (int) Math.round(h * look.scale());
|
||||
BufferedImage small = new BufferedImage(nw, nh, BufferedImage.TYPE_INT_RGB);
|
||||
Graphics2D sg = small.createGraphics();
|
||||
sg.setRenderingHint(RenderingHints.KEY_INTERPOLATION, RenderingHints.VALUE_INTERPOLATION_BILINEAR);
|
||||
sg.drawImage(out, 0, 0, nw, nh, null);
|
||||
sg.dispose();
|
||||
out = small;
|
||||
}
|
||||
if (look.noise() > 0) {
|
||||
Random rnd = new Random(42);
|
||||
for (int x = 0; x < out.getWidth(); x++) {
|
||||
for (int yy = 0; yy < out.getHeight(); yy++) {
|
||||
if (rnd.nextDouble() < look.noise()) {
|
||||
out.setRGB(x, yy, rnd.nextBoolean() ? 0x000000 : 0xFFFFFF);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
try (ByteArrayOutputStream bos = new ByteArrayOutputStream()) {
|
||||
ImageIO.write(out, "png", bos);
|
||||
return bos.toByteArray();
|
||||
} catch (IOException e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,68 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
import java.util.Set;
|
||||
|
||||
/**
|
||||
* Checks that an extracted receipt is arithmetically consistent. A model (or an OCR engine) can
|
||||
* misread a digit and still return perfectly valid JSON; the only way to notice is that the numbers
|
||||
* stop adding up. Every rule here is a sum the receipt itself promises.
|
||||
*/
|
||||
public final class ReceiptValidator {
|
||||
|
||||
private static final Set<String> CURRENCIES = Set.of("USD", "EUR", "GBP", "INR");
|
||||
|
||||
private ReceiptValidator() {
|
||||
}
|
||||
|
||||
public static List<String> problems(Receipt r) {
|
||||
List<String> out = new ArrayList<>();
|
||||
if (r == null) {
|
||||
return List.of("no receipt");
|
||||
}
|
||||
if (r.merchant() == null || r.merchant().isBlank()) {
|
||||
out.add("merchant missing");
|
||||
}
|
||||
if (r.date() == null) {
|
||||
out.add("date missing");
|
||||
}
|
||||
if (r.currency() == null || !CURRENCIES.contains(r.currency())) {
|
||||
out.add("currency " + r.currency() + " is not one of " + new java.util.TreeSet<>(CURRENCIES));
|
||||
}
|
||||
if (r.items() == null || r.items().isEmpty()) {
|
||||
out.add("no line items");
|
||||
return out;
|
||||
}
|
||||
BigDecimal sum = BigDecimal.ZERO;
|
||||
for (Receipt.Line l : r.items()) {
|
||||
if (l.unitPrice() == null || l.lineTotal() == null) {
|
||||
out.add("line '" + l.description() + "' has no price");
|
||||
continue;
|
||||
}
|
||||
BigDecimal expected = l.unitPrice().multiply(BigDecimal.valueOf(l.quantity()));
|
||||
if (expected.compareTo(l.lineTotal()) != 0) {
|
||||
out.add("line '" + l.description() + "': " + l.quantity() + " x " + l.unitPrice() + " is "
|
||||
+ expected + ", not " + l.lineTotal());
|
||||
}
|
||||
sum = sum.add(l.lineTotal());
|
||||
}
|
||||
if (r.subtotal() == null || sum.compareTo(r.subtotal()) != 0) {
|
||||
out.add("lines add up to " + sum + ", subtotal says " + r.subtotal());
|
||||
}
|
||||
if (r.subtotal() != null && r.tax() != null && r.total() != null
|
||||
&& r.subtotal().add(r.tax()).compareTo(r.total()) != 0) {
|
||||
out.add("subtotal " + r.subtotal() + " + tax " + r.tax() + " is " + r.subtotal().add(r.tax())
|
||||
+ ", total says " + r.total());
|
||||
}
|
||||
if (r.tax() == null || r.total() == null) {
|
||||
out.add("tax or total missing");
|
||||
}
|
||||
return out;
|
||||
}
|
||||
|
||||
public static boolean valid(Receipt r) {
|
||||
return problems(r).isEmpty();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,91 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.util.List;
|
||||
import java.util.Objects;
|
||||
|
||||
import com.ankurm.multimodal.ReceiptImages.Look;
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
/**
|
||||
* Ten receipts under four photo conditions through the whole pipeline. The "model" is OCR plus a
|
||||
* parser, so these numbers describe THIS backend and the plumbing around it, not any LLM.
|
||||
*/
|
||||
class AccuracyTest {
|
||||
|
||||
record Score(int fieldsRight, int fieldsTotal, boolean exact) {
|
||||
}
|
||||
|
||||
static Score score(Receipt truth, Receipt got) {
|
||||
int total = 6 + 4 * truth.items().size();
|
||||
if (got == null) {
|
||||
return new Score(0, total, false);
|
||||
}
|
||||
int ok = 0;
|
||||
ok += Objects.equals(truth.merchant(), got.merchant()) ? 1 : 0;
|
||||
ok += Objects.equals(truth.date(), got.date()) ? 1 : 0;
|
||||
ok += Objects.equals(truth.currency(), got.currency()) ? 1 : 0;
|
||||
ok += same(truth.subtotal(), got.subtotal()) ? 1 : 0;
|
||||
ok += same(truth.tax(), got.tax()) ? 1 : 0;
|
||||
ok += same(truth.total(), got.total()) ? 1 : 0;
|
||||
for (int i = 0; i < truth.items().size(); i++) {
|
||||
if (got.items() == null || i >= got.items().size()) {
|
||||
continue;
|
||||
}
|
||||
Receipt.Line a = truth.items().get(i);
|
||||
Receipt.Line b = got.items().get(i);
|
||||
ok += Objects.equals(a.description(), b.description()) ? 1 : 0;
|
||||
ok += a.quantity() == b.quantity() ? 1 : 0;
|
||||
ok += same(a.unitPrice(), b.unitPrice()) ? 1 : 0;
|
||||
ok += same(a.lineTotal(), b.lineTotal()) ? 1 : 0;
|
||||
}
|
||||
boolean exact = ok == total && got.items().size() == truth.items().size();
|
||||
return new Score(ok, total, exact);
|
||||
}
|
||||
|
||||
private static boolean same(java.math.BigDecimal a, java.math.BigDecimal b) {
|
||||
return a != null && b != null && a.compareTo(b) == 0;
|
||||
}
|
||||
|
||||
@Test
|
||||
void accuracyByPhotoCondition() throws Exception {
|
||||
List<Look> looks = List.of(Look.CLEAN, Look.TILTED, Look.SMALL, Look.NOISY);
|
||||
try (FakeVisionServer server = new FakeVisionServer(); Transcript t = new Transcript("05-accuracy.txt",
|
||||
"Ten receipts x four photo conditions through the real pipeline (OCR backend, NOT an LLM)")) {
|
||||
ReceiptExtractor ex = Wire.extractor(Models.openai(server.url()));
|
||||
t.line("%-20s %14s %14s %16s %20s", "condition", "fields right", "exact receipts", "pass validation",
|
||||
"passed but WRONG");
|
||||
int[] silentByLook = new int[looks.size()];
|
||||
for (int li = 0; li < looks.size(); li++) {
|
||||
Look look = looks.get(li);
|
||||
int right = 0, fields = 0, exact = 0, passed = 0, silent = 0;
|
||||
for (Receipt truth : Fixtures.all()) {
|
||||
Receipt got;
|
||||
try {
|
||||
got = ex.extractOnce(ReceiptImages.png(truth, look), MimeTypeUtils.IMAGE_PNG);
|
||||
} catch (RuntimeException e) {
|
||||
got = null;
|
||||
}
|
||||
Score s = score(truth, got);
|
||||
right += s.fieldsRight();
|
||||
fields += s.fieldsTotal();
|
||||
exact += s.exact() ? 1 : 0;
|
||||
boolean valid = got != null && ReceiptValidator.valid(got);
|
||||
passed += valid ? 1 : 0;
|
||||
silent += valid && !s.exact() ? 1 : 0;
|
||||
}
|
||||
silentByLook[li] = silent;
|
||||
t.line("%-20s %14s %14s %16s %20s", look.name(), right + "/" + fields + " (" + Math.round(100.0 * right / fields) + "%)",
|
||||
exact + "/10", passed + "/10", silent + "/10");
|
||||
if (look == Look.CLEAN) {
|
||||
assertThat(exact).isEqualTo(10);
|
||||
assertThat(silent).isZero();
|
||||
}
|
||||
}
|
||||
t.blank().line("'passed but WRONG' = the arithmetic checks were satisfied and the record still differs from the receipt.");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,39 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.awt.image.BufferedImage;
|
||||
import java.io.ByteArrayInputStream;
|
||||
import java.util.List;
|
||||
import javax.imageio.ImageIO;
|
||||
|
||||
import com.ankurm.multimodal.ImageTokens.Tier;
|
||||
import com.ankurm.multimodal.ReceiptImages.Look;
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
/** Size on the wire, and the estimated token cost from Anthropic's documented rule. */
|
||||
class CostTest {
|
||||
|
||||
@Test
|
||||
void sizeAndEstimate() throws Exception {
|
||||
Receipt r = Fixtures.all().get(1);
|
||||
try (Transcript t = new Transcript("07-cost.txt", "What an image costs: bytes now, tokens by documented formula")) {
|
||||
t.line("%-20s %9s %10s %12s %10s %12s", "look", "pixels", "png bytes", "base64 bytes", "tokens std", "tokens hi-res");
|
||||
for (Look look : List.of(Look.CLEAN, Look.SMALL)) {
|
||||
byte[] png = ReceiptImages.png(r, look);
|
||||
BufferedImage img = ImageIO.read(new ByteArrayInputStream(png));
|
||||
int b64 = java.util.Base64.getEncoder().encode(png).length;
|
||||
t.line("%-20s %9s %10d %12d %10d %12d", look.name(), img.getWidth() + "x" + img.getHeight(), png.length, b64,
|
||||
ImageTokens.anthropic(img.getWidth(), img.getHeight(), Tier.STANDARD),
|
||||
ImageTokens.anthropic(img.getWidth(), img.getHeight(), Tier.HIGH_RES));
|
||||
assertThat(b64).isBetween((int) (png.length * 1.33), (int) (png.length * 1.34) + 4);
|
||||
}
|
||||
t.blank().line("%-20s %9s %10s %12s %10s %12s", "phone photo (given)", "4032x3024", "-", "-",
|
||||
ImageTokens.anthropic(4032, 3024, Tier.STANDARD), ImageTokens.anthropic(4032, 3024, Tier.HIGH_RES));
|
||||
t.line("%-20s %9s %10s %12s %10s %12s", "1000x1000 (docs)", "1000x1000", "-", "-",
|
||||
ImageTokens.anthropic(1000, 1000, Tier.STANDARD), ImageTokens.anthropic(1000, 1000, Tier.HIGH_RES));
|
||||
assertThat(ImageTokens.patches(1000, 1000)).isEqualTo(1296);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,32 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
/** The happy path: image in, validated record out. Shows what the stand-in vision backend actually read. */
|
||||
class ExtractionTest {
|
||||
|
||||
@Test
|
||||
void cleanReceiptBecomesARecord() throws Exception {
|
||||
Receipt truth = Fixtures.all().get(0);
|
||||
byte[] png = ReceiptImages.png(truth, ReceiptImages.Look.CLEAN);
|
||||
try (FakeVisionServer server = new FakeVisionServer(); Transcript t = new Transcript("02-extraction.txt",
|
||||
"A clean receipt through ChatClient.entity(Receipt.class)")) {
|
||||
ReceiptExtractor.Result result = Wire.extractor(Models.openai(server.url())).extract(png);
|
||||
t.line("what the stand-in backend read from the pixels (tesseract OCR):");
|
||||
for (String l : server.lastOcr().strip().split("\\R")) {
|
||||
t.line(" | %s", l);
|
||||
}
|
||||
t.blank().line("the record Spring AI built from the model's JSON:");
|
||||
t.line(" %s", result.receipt());
|
||||
t.line("validation problems: %s; attempts used: %d", result.problems().isEmpty() ? "none" : result.problems(),
|
||||
result.attempts());
|
||||
assertThat(result.receipt()).isEqualTo(truth);
|
||||
assertThat(result.ok()).isTrue();
|
||||
assertThat(result.attempts()).isEqualTo(1);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,48 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
/** What the validate-and-retry loop does when the first answer is wrong, and when it cannot help. */
|
||||
class RepairLoopTest {
|
||||
|
||||
@Test
|
||||
void retryFixesARandomMisreadButNotADeterministicOne() throws Exception {
|
||||
try (FakeVisionServer server = new FakeVisionServer().misreadFirstAttempt(true);
|
||||
Transcript t = new Transcript("04-repair-loop.txt", "Validate, then ask once more with the problems quoted")) {
|
||||
Receipt truth = Fixtures.all().get(0);
|
||||
byte[] png = ReceiptImages.png(truth, ReceiptImages.Look.CLEAN);
|
||||
ReceiptExtractor.Result r = Wire.extractor(Models.openai(server.url())).extract(png);
|
||||
t.line("SIMULATED misread (the stand-in adds 0.10 to the total on the first attempt only):");
|
||||
t.line(" attempts: %d, problems after: %s, total: %s", r.attempts(), r.problems().isEmpty() ? "none" : r.problems(),
|
||||
r.receipt().total());
|
||||
String retryPrompt = server.seen().getLast().promptText();
|
||||
t.line(" text sent with the second attempt:");
|
||||
String[] promptLines = retryPrompt.strip().split("\\R");
|
||||
for (String l : promptLines) {
|
||||
if (l.contains("previous attempt")) {
|
||||
t.line(" | %s", l);
|
||||
}
|
||||
}
|
||||
t.line(" | (+ %d more lines: the system prompt and the JSON-schema format instructions Spring AI appends for entity())", promptLines.length - 1);
|
||||
assertThat(r.attempts()).isEqualTo(2);
|
||||
assertThat(r.ok()).isTrue();
|
||||
assertThat(r.receipt()).isEqualTo(truth);
|
||||
assertThat(server.seen()).hasSize(2);
|
||||
assertThat(server.seen().get(0).sha256()).isEqualTo(server.seen().get(1).sha256());
|
||||
}
|
||||
try (FakeVisionServer server = new FakeVisionServer(); Transcript t = new Transcript("04b-unreadable.txt",
|
||||
"Retrying an image the backend cannot read")) {
|
||||
Receipt truth = Fixtures.all().get(0);
|
||||
byte[] png = ReceiptImages.png(truth, ReceiptImages.Look.NOISY);
|
||||
ReceiptExtractor.Result r = Wire.extractor(Models.openai(server.url())).extract(png);
|
||||
t.line("noisy image, no simulation: attempts %d, requests made %d", r.attempts(), server.seen().size());
|
||||
t.line("problems left: %s", r.problems());
|
||||
t.line("The second attempt saw the same pixels and made the same mistake: a retry only helps when the error is not deterministic.");
|
||||
assertThat(r.attempts()).isEqualTo(2);
|
||||
assertThat(r.ok()).isFalse();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,43 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.net.URI;
|
||||
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import com.ankurm.multimodal.support.Wire.Provider;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
/** What Media(type, URI) does on each real model: send a link, or fetch and embed the bytes. */
|
||||
class UrlVsBytesTest {
|
||||
|
||||
@Test
|
||||
void imageByUrl() throws Exception {
|
||||
try (FakeVisionServer server = new FakeVisionServer(); Transcript t = new Transcript("06-url-vs-bytes.txt",
|
||||
"Media given as a URI instead of bytes")) {
|
||||
URI uri = URI.create("https://example.invalid/receipts/2026-03-14.png");
|
||||
for (Provider p : Wire.providers(server.url())) {
|
||||
int before = server.seen().size();
|
||||
String outcome;
|
||||
try {
|
||||
ChatClient.create(p.model()).prompt().user(u -> u.text("Extract this receipt.").media(org.springframework.ai.content.Media.builder().mimeType(MimeTypeUtils.IMAGE_PNG).data(uri).build()))
|
||||
.call().content();
|
||||
FakeVisionServer.Seen s = server.seen().get(before);
|
||||
outcome = "request sent; image part carries: " + s.mimeAsSent();
|
||||
} catch (RuntimeException e) {
|
||||
outcome = "failed before sending: " + e.getClass().getSimpleName() + ": " + firstLine(e);
|
||||
}
|
||||
t.line("%-9s %s", p.name(), outcome);
|
||||
}
|
||||
t.blank().line("(example.invalid cannot resolve, so a provider that fetches the URL itself would fail here; the fake server never fetches.)");
|
||||
assertThat(server.seen().size()).isGreaterThanOrEqualTo(0);
|
||||
}
|
||||
}
|
||||
|
||||
private static String firstLine(Throwable e) {
|
||||
String m = String.valueOf(e.getMessage());
|
||||
return m.lines().findFirst().orElse("").strip();
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,49 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.util.ArrayList;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import org.junit.jupiter.api.Test;
|
||||
|
||||
/** One deliberately wrong receipt per rule, and one error the rules cannot see. */
|
||||
class ValidatorTest {
|
||||
|
||||
private final Receipt good = Fixtures.all().get(0);
|
||||
|
||||
private Receipt with(String merchant, String currency, List<Receipt.Line> items, String sub, String tax, String total) {
|
||||
return new Receipt(merchant, good.date(), currency, items, new BigDecimal(sub), new BigDecimal(tax), new BigDecimal(total));
|
||||
}
|
||||
|
||||
@Test
|
||||
void eachRuleCatchesItsError() {
|
||||
List<Receipt.Line> lines = good.items();
|
||||
List<Receipt.Line> badLine = new ArrayList<>(lines);
|
||||
badLine.set(0, new Receipt.Line("Flat White", 2, new BigDecimal("4.50"), new BigDecimal("9.50")));
|
||||
List<Receipt.Line> misread = new ArrayList<>(lines);
|
||||
misread.set(1, new Receipt.Line("Croissant", 1, new BigDecimal("3.15"), new BigDecimal("3.15")));
|
||||
List<Receipt.Line> typo = new ArrayList<>(lines);
|
||||
typo.set(1, new Receipt.Line("Croissent", 1, new BigDecimal("3.75"), new BigDecimal("3.75")));
|
||||
try (Transcript t = new Transcript("03-validator.txt", "What the arithmetic checks catch, and what they cannot")) {
|
||||
record Case(String name, Receipt r, boolean shouldBeCaught) {
|
||||
}
|
||||
List<Case> cases = List.of(
|
||||
new Case("correct receipt", good, false),
|
||||
new Case("line total is not qty x unit price", with("CORNER CAFE", "USD", badLine, "13.25", "1.12", "14.37"), true),
|
||||
new Case("a digit misread in one unit price (3.75 -> 3.15)", with("CORNER CAFE", "USD", misread, "12.75", "1.12", "13.87"), true),
|
||||
new Case("total does not equal subtotal + tax", with("CORNER CAFE", "USD", lines, "12.75", "1.12", "13.97"), true),
|
||||
new Case("currency the app does not support", with("CORNER CAFE", "US$", lines, "12.75", "1.12", "13.87"), true),
|
||||
new Case("merchant name typo (text, not arithmetic)", with("CORNER CAFF", "USD", lines, "12.75", "1.12", "13.87"), false),
|
||||
new Case("item name typo (text, not arithmetic)", with("CORNER CAFE", "USD", typo, "12.75", "1.12", "13.87"), false));
|
||||
for (Case c : cases) {
|
||||
List<String> problems = ReceiptValidator.problems(c.r());
|
||||
t.line("%-52s -> %s", c.name(), problems.isEmpty() ? "PASSES validation" : problems);
|
||||
assertThat(!problems.isEmpty()).as(c.name()).isEqualTo(c.shouldBeCaught());
|
||||
}
|
||||
t.blank().line("The last two are wrong and pass: arithmetic cannot check spelling.");
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,45 @@
|
||||
package com.ankurm.multimodal;
|
||||
|
||||
import static org.assertj.core.api.Assertions.assertThat;
|
||||
|
||||
import com.ankurm.multimodal.support.*;
|
||||
import com.ankurm.multimodal.support.Wire.Provider;
|
||||
import org.junit.jupiter.api.Test;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
import tools.jackson.databind.json.JsonMapper;
|
||||
|
||||
/** What each real Spring AI model puts on the wire for the SAME image, and that the bytes arrive intact. */
|
||||
class WireFormatTest {
|
||||
|
||||
private static final JsonMapper JSON = JsonMapper.builder().build();
|
||||
|
||||
@Test
|
||||
void sameImageThreeWireFormats() throws Exception {
|
||||
byte[] png = ReceiptImages.png(Fixtures.all().get(0), ReceiptImages.Look.CLEAN);
|
||||
String sent = FakeVisionServerHashes.sha(png);
|
||||
try (FakeVisionServer server = new FakeVisionServer(); Transcript t = new Transcript("01-wire-formats.txt",
|
||||
"The same receipt image through three real Spring AI chat models (" + png.length + " bytes, PNG)")) {
|
||||
for (Provider p : Wire.providers(server.url())) {
|
||||
Receipt r = Wire.extractor(p.model()).extractOnce(png, MimeTypeUtils.IMAGE_PNG);
|
||||
assertThat(r.total()).isEqualByComparingTo("13.87");
|
||||
FakeVisionServer.Seen seen = server.seen().getLast();
|
||||
JsonNode req = JSON.readTree(seen.requestJson());
|
||||
JsonNode msg = switch (p.name()) {
|
||||
case "openai" -> req.path("messages").get(1).path("content");
|
||||
case "anthropic" -> req.path("messages").get(0).path("content");
|
||||
default -> req.path("messages").get(1);
|
||||
};
|
||||
t.line("%-9s image as sent:", p.name());
|
||||
t.line(" %s", Wire.brief(msg));
|
||||
t.line(" mime label: %s | bytes the server decoded: %d | sha256 matches: %s", seen.mimeAsSent(),
|
||||
seen.bytes(), seen.sha256().equals(sent));
|
||||
assertThat(seen.sha256()).isEqualTo(sent);
|
||||
assertThat(seen.bytes()).isEqualTo(png.length);
|
||||
}
|
||||
long distinct = server.seen().stream().map(FakeVisionServer.Seen::sha256).distinct().count();
|
||||
t.blank().line("distinct image hashes seen by the server across the three providers: %d", distinct);
|
||||
assertThat(distinct).isEqualTo(1);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,267 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
import java.awt.image.BufferedImage;
|
||||
import java.io.ByteArrayInputStream;
|
||||
import java.io.IOException;
|
||||
import java.io.OutputStream;
|
||||
import java.net.InetSocketAddress;
|
||||
import java.nio.charset.StandardCharsets;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
import java.security.MessageDigest;
|
||||
import java.util.ArrayList;
|
||||
import java.util.Base64;
|
||||
import java.util.HexFormat;
|
||||
import java.util.List;
|
||||
import java.util.concurrent.CopyOnWriteArrayList;
|
||||
import java.util.concurrent.TimeUnit;
|
||||
import java.util.regex.Matcher;
|
||||
import java.util.regex.Pattern;
|
||||
import javax.imageio.ImageIO;
|
||||
|
||||
import com.sun.net.httpserver.HttpExchange;
|
||||
import com.sun.net.httpserver.HttpServer;
|
||||
import tools.jackson.databind.JsonNode;
|
||||
import tools.jackson.databind.json.JsonMapper;
|
||||
import tools.jackson.databind.node.ArrayNode;
|
||||
import tools.jackson.databind.node.ObjectNode;
|
||||
|
||||
/**
|
||||
* A stand-in for three vision APIs on one port: OpenAI chat completions, Anthropic messages and
|
||||
* Ollama chat. The REAL Spring AI model classes talk to it, so the request each one builds is the
|
||||
* request it would send to the real service.
|
||||
*
|
||||
* <p>It is not a language model. It pulls the image out of whatever request shape arrived, decodes
|
||||
* it, runs the {@code tesseract} OCR program on the pixels, and parses the text with a few regular
|
||||
* expressions into the receipt JSON. So it genuinely depends on what is in the image, but its
|
||||
* accuracy is the accuracy of an OCR engine plus a hand-written parser, NOT of any vision model.
|
||||
* Treat every accuracy number it produces as a property of this pipeline's plumbing and of OCR,
|
||||
* never as a statement about GPT, Claude or Gemini.
|
||||
*/
|
||||
public class FakeVisionServer implements AutoCloseable {
|
||||
|
||||
/** What the server actually received for one request. */
|
||||
public record Seen(String provider, String mimeAsSent, int bytes, String sha256, int width, int height,
|
||||
String requestJson, String promptText) {
|
||||
}
|
||||
|
||||
private static final JsonMapper JSON = JsonMapper.builder().build();
|
||||
|
||||
private final HttpServer server;
|
||||
|
||||
private final List<Seen> seen = new CopyOnWriteArrayList<>();
|
||||
|
||||
private volatile String lastOcr = "";
|
||||
|
||||
private volatile boolean misreadFirstAttempt;
|
||||
|
||||
/** Simulation switch: add 0.10 to the total unless the prompt says a previous attempt had problems. */
|
||||
public FakeVisionServer misreadFirstAttempt(boolean on) {
|
||||
this.misreadFirstAttempt = on;
|
||||
return this;
|
||||
}
|
||||
|
||||
public FakeVisionServer() throws IOException {
|
||||
server = HttpServer.create(new InetSocketAddress("127.0.0.1", 0), 0);
|
||||
server.createContext("/v1/chat/completions", ex -> handle(ex, "openai"));
|
||||
server.createContext("/v1/messages", ex -> handle(ex, "anthropic"));
|
||||
server.createContext("/api/chat", ex -> handle(ex, "ollama"));
|
||||
server.start();
|
||||
}
|
||||
|
||||
public String url() {
|
||||
return "http://127.0.0.1:" + server.getAddress().getPort();
|
||||
}
|
||||
|
||||
public List<Seen> seen() {
|
||||
return seen;
|
||||
}
|
||||
|
||||
public String lastOcr() {
|
||||
return lastOcr;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
server.stop(0);
|
||||
}
|
||||
|
||||
private void handle(HttpExchange ex, String provider) throws IOException {
|
||||
String body = new String(ex.getRequestBody().readAllBytes(), StandardCharsets.UTF_8);
|
||||
JsonNode req = JSON.readTree(body);
|
||||
String mime = "?";
|
||||
String b64 = null;
|
||||
StringBuilder prompt = new StringBuilder();
|
||||
for (JsonNode m : req.path("messages")) {
|
||||
if (provider.equals("ollama")) {
|
||||
for (JsonNode img : m.path("images")) {
|
||||
b64 = img.asString();
|
||||
mime = "(none: Ollama sends bare base64)";
|
||||
}
|
||||
if (m.path("content").isString()) {
|
||||
prompt.append(m.path("content").asString()).append('\n');
|
||||
}
|
||||
continue;
|
||||
}
|
||||
JsonNode content = m.path("content");
|
||||
if (content.isString()) {
|
||||
prompt.append(content.asString()).append('\n');
|
||||
continue;
|
||||
}
|
||||
for (JsonNode part : content) {
|
||||
String type = part.path("type").asString();
|
||||
if (type.equals("text")) {
|
||||
prompt.append(part.path("text").asString()).append('\n');
|
||||
} else if (type.equals("image_url")) {
|
||||
String url = part.path("image_url").path("url").asString();
|
||||
Matcher dm = Pattern.compile("^data:([^;]+);base64,(.*)$", Pattern.DOTALL).matcher(url);
|
||||
if (dm.matches()) {
|
||||
mime = dm.group(1);
|
||||
b64 = dm.group(2);
|
||||
} else {
|
||||
mime = "URL: " + url;
|
||||
}
|
||||
} else if (type.equals("image")) {
|
||||
if (part.path("source").path("type").asString().equals("url")) {
|
||||
mime = "URL: " + part.path("source").path("url").asString();
|
||||
} else {
|
||||
mime = part.path("source").path("media_type").asString();
|
||||
b64 = part.path("source").path("data").asString();
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
String answer;
|
||||
if (b64 == null) {
|
||||
answer = "{}";
|
||||
seen.add(new Seen(provider, mime, 0, "", 0, 0, body, prompt.toString()));
|
||||
} else {
|
||||
byte[] png;
|
||||
try {
|
||||
png = Base64.getDecoder().decode(b64);
|
||||
} catch (IllegalArgumentException notBase64) {
|
||||
png = null;
|
||||
}
|
||||
BufferedImage img = png == null ? null : ImageIO.read(new ByteArrayInputStream(png));
|
||||
if (img == null) {
|
||||
String shown = b64.length() > 70 ? b64.substring(0, 70) + "..." : b64;
|
||||
seen.add(new Seen(provider, "NOT AN IMAGE: " + shown, 0, "", 0, 0, body, prompt.toString()));
|
||||
send(ex, provider, "{}");
|
||||
return;
|
||||
}
|
||||
seen.add(new Seen(provider, mime, png.length, sha256(png), img.getWidth(), img.getHeight(), body,
|
||||
prompt.toString()));
|
||||
lastOcr = ocr(png);
|
||||
answer = parse(lastOcr);
|
||||
if (misreadFirstAttempt && !prompt.toString().contains("previous attempt")) {
|
||||
ObjectNode n = (ObjectNode) JSON.readTree(answer);
|
||||
n.put("total", n.path("total").decimalValue().add(new java.math.BigDecimal("0.10")));
|
||||
answer = JSON.writeValueAsString(n);
|
||||
}
|
||||
}
|
||||
send(ex, provider, answer);
|
||||
}
|
||||
|
||||
private void send(HttpExchange ex, String provider, String answer) throws IOException {
|
||||
String reply = switch (provider) {
|
||||
case "openai" -> "{\"id\":\"chatcmpl-fake\",\"object\":\"chat.completion\",\"created\":1700000000,"
|
||||
+ "\"model\":\"gpt-4o-mini-2024-07-18\",\"choices\":[{\"index\":0,\"message\":{\"role\":\"assistant\","
|
||||
+ "\"content\":" + JSON.writeValueAsString(answer) + "},\"finish_reason\":\"stop\"}],"
|
||||
+ "\"usage\":{\"prompt_tokens\":10,\"completion_tokens\":10,\"total_tokens\":20}}";
|
||||
case "anthropic" -> "{\"id\":\"msg_fake\",\"type\":\"message\",\"role\":\"assistant\","
|
||||
+ "\"model\":\"claude-sonnet-4-5\",\"content\":[{\"type\":\"text\",\"text\":"
|
||||
+ JSON.writeValueAsString(answer) + "}],\"stop_reason\":\"end_turn\",\"stop_sequence\":null,"
|
||||
+ "\"usage\":{\"input_tokens\":10,\"output_tokens\":10}}";
|
||||
default -> "{\"model\":\"llava\",\"created_at\":\"2026-01-01T00:00:00Z\",\"message\":{\"role\":\"assistant\","
|
||||
+ "\"content\":" + JSON.writeValueAsString(answer) + "},\"done\":true,\"done_reason\":\"stop\","
|
||||
+ "\"prompt_eval_count\":10,\"eval_count\":10}";
|
||||
};
|
||||
ex.getResponseHeaders().add("Content-Type", "application/json");
|
||||
byte[] out = reply.getBytes(StandardCharsets.UTF_8);
|
||||
ex.sendResponseHeaders(200, out.length);
|
||||
try (OutputStream os = ex.getResponseBody()) {
|
||||
os.write(out);
|
||||
}
|
||||
}
|
||||
|
||||
static String sha256(byte[] data) {
|
||||
try {
|
||||
return HexFormat.of().formatHex(MessageDigest.getInstance("SHA-256").digest(data));
|
||||
} catch (Exception e) {
|
||||
throw new IllegalStateException(e);
|
||||
}
|
||||
}
|
||||
|
||||
/** Runs tesseract on the image bytes. */
|
||||
public static String ocr(byte[] png) {
|
||||
try {
|
||||
Path f = Files.createTempFile("receipt", ".png");
|
||||
Files.write(f, png);
|
||||
Process p = new ProcessBuilder("tesseract", f.toString(), "stdout", "--psm", "6")
|
||||
.redirectError(ProcessBuilder.Redirect.DISCARD).start();
|
||||
String text = new String(p.getInputStream().readAllBytes(), StandardCharsets.UTF_8);
|
||||
p.waitFor(60, TimeUnit.SECONDS);
|
||||
Files.delete(f);
|
||||
return text;
|
||||
} catch (IOException | InterruptedException e) {
|
||||
throw new IllegalStateException("tesseract failed", e);
|
||||
}
|
||||
}
|
||||
|
||||
private static final Pattern DATE = Pattern.compile("(\\d{4}-\\d{2}-\\d{2})");
|
||||
private static final Pattern CUR = Pattern.compile("CURRENCY:\\s*([A-Z]{3})");
|
||||
private static final Pattern LINE = Pattern.compile("^(.+?)\\s+(\\d+)\\s*[xX]\\s*([\\d.,]+)\\s+([\\d.,]+)\\s*$");
|
||||
private static final Pattern SUB = Pattern.compile("^SUBTOTAL\\s+([\\d.,]+)");
|
||||
private static final Pattern TAX = Pattern.compile("^TAX\\s+([\\d.,]+)");
|
||||
private static final Pattern TOT = Pattern.compile("^TOTAL\\s+([\\d.,]+)");
|
||||
|
||||
/** Turns OCR text into the receipt JSON; fields it cannot read are null. */
|
||||
public static String parse(String text) {
|
||||
ObjectNode r = JSON.createObjectNode();
|
||||
r.putNull("merchant");
|
||||
r.putNull("date");
|
||||
r.putNull("currency");
|
||||
ArrayNode items = r.putArray("items");
|
||||
r.putNull("subtotal");
|
||||
r.putNull("tax");
|
||||
r.putNull("total");
|
||||
List<String> lines = new ArrayList<>();
|
||||
for (String s : text.split("\\R")) {
|
||||
if (!s.isBlank()) {
|
||||
lines.add(s.strip());
|
||||
}
|
||||
}
|
||||
if (!lines.isEmpty()) {
|
||||
r.put("merchant", lines.get(0));
|
||||
}
|
||||
for (String s : lines) {
|
||||
Matcher m;
|
||||
if ((m = DATE.matcher(s)).find()) {
|
||||
r.put("date", m.group(1));
|
||||
} else if ((m = CUR.matcher(s)).find()) {
|
||||
r.put("currency", m.group(1));
|
||||
} else if ((m = SUB.matcher(s)).find()) {
|
||||
put(r, "subtotal", m.group(1));
|
||||
} else if ((m = TAX.matcher(s)).find()) {
|
||||
put(r, "tax", m.group(1));
|
||||
} else if ((m = TOT.matcher(s)).find()) {
|
||||
put(r, "total", m.group(1));
|
||||
} else if ((m = LINE.matcher(s)).matches()) {
|
||||
ObjectNode it = items.addObject();
|
||||
it.put("description", m.group(1).strip());
|
||||
it.put("quantity", Integer.parseInt(m.group(2)));
|
||||
put(it, "unitPrice", m.group(3));
|
||||
put(it, "lineTotal", m.group(4));
|
||||
}
|
||||
}
|
||||
return JSON.writeValueAsString(r);
|
||||
}
|
||||
|
||||
private static void put(ObjectNode n, String field, String num) {
|
||||
try {
|
||||
n.put(field, new java.math.BigDecimal(num.replace(',', '.')));
|
||||
} catch (NumberFormatException e) {
|
||||
n.putNull(field);
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,11 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
public final class FakeVisionServerHashes {
|
||||
|
||||
private FakeVisionServerHashes() {
|
||||
}
|
||||
|
||||
public static String sha(byte[] data) {
|
||||
return FakeVisionServer.sha256(data);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
import java.math.BigDecimal;
|
||||
import java.time.LocalDate;
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.multimodal.Receipt;
|
||||
import com.ankurm.multimodal.Receipt.Line;
|
||||
|
||||
/** Ten ground-truth receipts, written by hand. Each is internally consistent. */
|
||||
public final class Fixtures {
|
||||
|
||||
private Fixtures() {
|
||||
}
|
||||
|
||||
private static BigDecimal d(String s) {
|
||||
return new BigDecimal(s);
|
||||
}
|
||||
|
||||
private static Line line(String name, int q, String unit) {
|
||||
BigDecimal u = d(unit);
|
||||
return new Line(name, q, u, u.multiply(BigDecimal.valueOf(q)));
|
||||
}
|
||||
|
||||
private static Receipt receipt(String merchant, String date, String cur, String tax, Line... lines) {
|
||||
BigDecimal sub = BigDecimal.ZERO;
|
||||
for (Line l : lines) {
|
||||
sub = sub.add(l.lineTotal());
|
||||
}
|
||||
BigDecimal t = d(tax);
|
||||
return new Receipt(merchant, LocalDate.parse(date), cur, List.of(lines), sub, t, sub.add(t));
|
||||
}
|
||||
|
||||
public static List<Receipt> all() {
|
||||
return List.of(
|
||||
receipt("CORNER CAFE", "2026-03-14", "USD", "1.12", line("Flat White", 2, "4.50"), line("Croissant", 1, "3.75")),
|
||||
receipt("GREEN GROCER", "2026-03-15", "USD", "2.40", line("Apples", 3, "1.20"), line("Bread", 2, "2.95"), line("Milk", 1, "1.89")),
|
||||
receipt("BOOK NOOK", "2026-03-16", "GBP", "4.00", line("Novel", 2, "9.99"), line("Bookmark", 4, "1.50")),
|
||||
receipt("PIZZA ROMA", "2026-03-17", "EUR", "3.20", line("Margherita", 2, "8.50"), line("Cola", 3, "2.50"), line("Tiramisu", 1, "5.00")),
|
||||
receipt("HARDWARE HUB", "2026-03-18", "USD", "6.18", line("Screws", 5, "2.10"), line("Hammer", 1, "14.99"), line("Tape", 2, "3.49")),
|
||||
receipt("CHAI POINT", "2026-03-19", "INR", "18.00", line("Masala Chai", 4, "40.00"), line("Samosa", 6, "15.00")),
|
||||
receipt("PET PLACE", "2026-03-20", "USD", "3.05", line("Dog Food", 1, "24.99"), line("Chew Toy", 2, "6.50")),
|
||||
receipt("FLOWER BAR", "2026-03-21", "EUR", "2.70", line("Tulips", 3, "7.00"), line("Vase", 1, "12.00")),
|
||||
receipt("TECH DEPOT", "2026-03-22", "USD", "7.59", line("USB Cable", 3, "9.99"), line("Mouse", 1, "19.50"), line("Hub", 1, "29.00")),
|
||||
receipt("NIGHT MARKET", "2026-03-23", "GBP", "1.45", line("Noodles", 2, "6.75"), line("Tea", 2, "2.20")));
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,42 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
import org.springframework.ai.anthropic.AnthropicChatModel;
|
||||
import org.springframework.ai.anthropic.AnthropicChatOptions;
|
||||
import org.springframework.ai.chat.client.ChatClient;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.ai.ollama.OllamaChatModel;
|
||||
import org.springframework.ai.ollama.api.OllamaApi;
|
||||
import org.springframework.ai.ollama.api.OllamaChatOptions;
|
||||
import org.springframework.ai.openai.OpenAiChatModel;
|
||||
import org.springframework.ai.openai.OpenAiChatOptions;
|
||||
|
||||
/** The three REAL Spring AI chat models, pointed at the fake server. */
|
||||
public final class Models {
|
||||
|
||||
private Models() {
|
||||
}
|
||||
|
||||
public static ChatModel openai(String baseUrl) {
|
||||
return OpenAiChatModel.builder()
|
||||
.options(OpenAiChatOptions.builder().baseUrl(baseUrl + "/v1").apiKey("test").model("gpt-4o-mini").build())
|
||||
.build();
|
||||
}
|
||||
|
||||
public static ChatModel anthropic(String baseUrl) {
|
||||
return AnthropicChatModel.builder()
|
||||
.options(AnthropicChatOptions.builder().baseUrl(baseUrl).apiKey("test").model("claude-sonnet-4-5")
|
||||
.maxTokens(1024).build())
|
||||
.build();
|
||||
}
|
||||
|
||||
public static ChatModel ollama(String baseUrl) {
|
||||
return OllamaChatModel.builder()
|
||||
.ollamaApi(OllamaApi.builder().baseUrl(baseUrl).build())
|
||||
.options(OllamaChatOptions.builder().model("llava").build())
|
||||
.build();
|
||||
}
|
||||
|
||||
public static ChatClient client(ChatModel m) {
|
||||
return ChatClient.create(m);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,47 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
import java.io.IOException;
|
||||
import java.io.PrintWriter;
|
||||
import java.io.StringWriter;
|
||||
import java.nio.file.Files;
|
||||
import java.nio.file.Path;
|
||||
|
||||
/**
|
||||
* Writes a numbered transcript under {@code output/} (repository root, not {@code docs/}) and
|
||||
* echoes it to the console. Every console block quoted in the article comes out of one of these
|
||||
* files verbatim.
|
||||
*/
|
||||
public final class Transcript implements AutoCloseable {
|
||||
|
||||
private final Path path;
|
||||
private final StringWriter buffer = new StringWriter();
|
||||
private final PrintWriter out = new PrintWriter(buffer);
|
||||
|
||||
public Transcript(String fileName, String title) {
|
||||
this.path = Path.of("output", fileName);
|
||||
out.println("# " + title);
|
||||
out.println();
|
||||
}
|
||||
|
||||
public Transcript line(String format, Object... args) {
|
||||
out.println(args.length == 0 ? format : String.format(format, args));
|
||||
return this;
|
||||
}
|
||||
|
||||
public Transcript blank() {
|
||||
out.println();
|
||||
return this;
|
||||
}
|
||||
|
||||
@Override
|
||||
public void close() {
|
||||
out.flush();
|
||||
try {
|
||||
Files.createDirectories(path.getParent());
|
||||
Files.writeString(path, buffer.toString());
|
||||
} catch (IOException e) {
|
||||
throw new IllegalStateException("could not write " + path, e);
|
||||
}
|
||||
System.out.print(buffer);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,73 @@
|
||||
package com.ankurm.multimodal.support;
|
||||
|
||||
import java.util.List;
|
||||
|
||||
import com.ankurm.multimodal.ReceiptExtractor;
|
||||
import org.springframework.ai.chat.model.ChatModel;
|
||||
import org.springframework.util.MimeTypeUtils;
|
||||
|
||||
/** Shared test helpers. */
|
||||
public final class Wire {
|
||||
|
||||
public record Provider(String name, ChatModel model) {
|
||||
}
|
||||
|
||||
private Wire() {
|
||||
}
|
||||
|
||||
public static List<Provider> providers(String url) {
|
||||
return List.of(new Provider("openai", Models.openai(url)), new Provider("anthropic", Models.anthropic(url)),
|
||||
new Provider("ollama", Models.ollama(url)));
|
||||
}
|
||||
|
||||
public static ReceiptExtractor extractor(ChatModel m) {
|
||||
return new ReceiptExtractor(Models.client(m));
|
||||
}
|
||||
|
||||
/** Replaces any long base64 run with a short placeholder so a request can be printed. */
|
||||
public static String elide(String json) {
|
||||
java.util.regex.Matcher m = java.util.regex.Pattern.compile("[A-Za-z0-9+/=]{200,}").matcher(json);
|
||||
StringBuilder sb = new StringBuilder();
|
||||
while (m.find()) {
|
||||
m.appendReplacement(sb, "<" + m.group().length() + " base64 characters>");
|
||||
}
|
||||
m.appendTail(sb);
|
||||
return sb.toString();
|
||||
}
|
||||
|
||||
/** Shortens every long string inside a JSON value (prompt text, base64) so a request fits on a line. */
|
||||
public static String brief(tools.jackson.databind.JsonNode n) {
|
||||
tools.jackson.databind.json.JsonMapper m = tools.jackson.databind.json.JsonMapper.builder().build();
|
||||
return briefNode(m, n).toString();
|
||||
}
|
||||
|
||||
private static tools.jackson.databind.JsonNode briefNode(tools.jackson.databind.json.JsonMapper m, tools.jackson.databind.JsonNode n) {
|
||||
if (n.isString()) {
|
||||
String s = n.asString();
|
||||
if (s.matches("[A-Za-z0-9+/=]{200,}")) {
|
||||
return m.getNodeFactory().stringNode("<" + s.length() + " base64 characters>");
|
||||
}
|
||||
java.util.regex.Matcher dm = java.util.regex.Pattern.compile("(?s)data:([^;]+);base64,([A-Za-z0-9+/=]+)").matcher(s);
|
||||
String cut = dm.matches() ? "data:" + dm.group(1) + ";base64,<" + dm.group(2).length() + " base64 characters>" : s;
|
||||
if (dm.matches()) {
|
||||
return m.getNodeFactory().stringNode(cut);
|
||||
}
|
||||
return m.getNodeFactory().stringNode(cut.length() > 60 ? cut.substring(0, 40).replace("\n", " ") + "... (" + cut.length() + " characters)" : cut);
|
||||
}
|
||||
if (n.isObject()) {
|
||||
tools.jackson.databind.node.ObjectNode o = m.createObjectNode();
|
||||
n.properties().forEach(e -> o.set(e.getKey(), briefNode(m, e.getValue())));
|
||||
return o;
|
||||
}
|
||||
if (n.isArray()) {
|
||||
tools.jackson.databind.node.ArrayNode a = m.createArrayNode();
|
||||
n.forEach(x -> a.add(briefNode(m, x)));
|
||||
return a;
|
||||
}
|
||||
return n;
|
||||
}
|
||||
|
||||
public static String pngMime() {
|
||||
return MimeTypeUtils.IMAGE_PNG_VALUE;
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user