package classification import ( "context" "errors" "io" "net/http" "reflect" "strings" "testing" "time" "finance-duck/internal/domain" "finance-duck/internal/ratelimit" ) func batchRows() (domain.Facts, domain.Facts, domain.Dataset) { f1, d := fixture() f1.Counterparty = "Coffee House" f2 := f1 f2.ID, f2.Fingerprint, f2.ExternalID = "tx_two", "fp_two", "ext_two" f2.Amount = "-4.30" f2.Counterparty = "Kleins Backstube" return f1, f2, d } // One request classifies every row: the prompt carries all transactions with // refs, and each answer resolves independently against the registry. func TestBatchClassifiesEveryRowInOneRequest(t *testing.T) { f1, f2, d := batchRows() calls := 0 c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { calls++ prompt := decodeClassificationPrompt(t, r) if len(prompt.Transactions) != 2 { t.Errorf("batch prompt missing transactions: %+v", prompt.Transactions) w.WriteHeader(http.StatusBadRequest) return } category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food"))) merchant, tag := "", "" for _, candidate := range prompt.Merchants { if candidate.Name == "coffee house" { merchant = candidate.ID } } for _, candidate := range prompt.Tags { if candidate.Name == "daily" { tag = candidate.ID } } if merchant == "" || tag == "" { t.Error("batch prompt lost Coffee House or Daily") } for _, row := range prompt.Transactions { if row.Amount == "" || row.Currency != "EUR" { t.Errorf("row %s lost amount or currency: %+v", row.Ref, row) } } reply(w, `{"transactions":[{"ref":"`+prompt.Transactions[1].Ref+`","merchant_id":null,"new_merchant":"Kleins Backstube","category_id":"`+category+`","tag_ids":[],"confidence":"medium"},`+ `{"ref":"`+prompt.Transactions[0].Ref+`","merchant_id":"`+merchant+`","new_merchant":null,"category_id":"`+category+`","tag_ids":["`+tag+`"],"confidence":"high"}]}`) }) results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d) if calls != 1 { t.Fatalf("expected one provider request for the batch, got %d", calls) } if results[0].Err != nil || results[1].Err != nil { t.Fatalf("batch rows failed: %v %v", results[0].Err, results[1].Err) } first := results[0].Proposal.Enrichment if first.MerchantID != "mer_coffee" || first.CategoryID != "cat_food" || !reflect.DeepEqual(first.TagIDs, []string{"tag_daily"}) || first.Classification.Confidence != "high" { t.Fatalf("first row lost: %+v", first) } second := results[1].Proposal if second.NewMerchant == nil || second.NewMerchant.Name != "Kleins Backstube" || !reflect.DeepEqual(second.NewMerchant.Aliases, []string{"Kleins Backstube"}) || second.Enrichment.MerchantID != second.NewMerchant.ID || second.Enrichment.CategoryID != "cat_food" || len(second.Enrichment.TagIDs) != 0 || second.Enrichment.Classification.Confidence != "medium" { t.Fatalf("second row lost: %+v", second) } } // One row's out-of-registry answer fails only that row. func TestBatchIsolatesInvalidRows(t *testing.T) { f1, f2, d := batchRows() c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { prompt := decodeClassificationPrompt(t, r) category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food"))) reply(w, `{"transactions":[{"ref":"`+prompt.Transactions[0].Ref+`","merchant_id":null,"new_merchant":null,"category_id":"`+category+`","tag_ids":[],"confidence":"high"},`+ `{"ref":"`+prompt.Transactions[1].Ref+`","merchant_id":null,"new_merchant":null,"category_id":"c999999","tag_ids":[],"confidence":"high"}]}`) }) results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d) if results[0].Err != nil || results[0].Proposal.Enrichment.CategoryID != "cat_food" { t.Fatalf("healthy row poisoned: %+v", results[0]) } if results[1].Err == nil || results[1].Proposal.Enrichment.Classification.Source != "fallback" { t.Fatalf("forged category accepted: %+v", results[1]) } } // Two rows naming the same new business share one minted merchant. func TestBatchSharesOneMintedMerchant(t *testing.T) { f1, f2, d := batchRows() c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { prompt := decodeClassificationPrompt(t, r) category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food"))) reply(w, `{"transactions":[{"ref":"`+prompt.Transactions[0].Ref+`","merchant_id":null,"new_merchant":"REWE","category_id":"`+category+`","tag_ids":[],"confidence":"high"},`+ `{"ref":"`+prompt.Transactions[1].Ref+`","merchant_id":null,"new_merchant":"REWE","category_id":"`+category+`","tag_ids":[],"confidence":"high"}]}`) }) results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d) if results[0].Err != nil || results[1].Err != nil { t.Fatalf("batch failed: %v %v", results[0].Err, results[1].Err) } a, b := results[0].Proposal, results[1].Proposal if a.NewMerchant == nil || b.NewMerchant == nil || a.NewMerchant.ID != b.NewMerchant.ID || a.Enrichment.MerchantID != b.Enrichment.MerchantID { t.Fatalf("duplicate merchants minted: %+v %+v", a.NewMerchant, b.NewMerchant) } } // A request-level rate limit fails every row and arms the shared cooldown. func TestBatchRateLimitFailsAllRowsAndArmsCooldown(t *testing.T) { f1, f2, d := batchRows() calls := 0 c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { calls++ _, _ = io.WriteString(w, `{"error":{"code":429,"message":"private"},"choices":[]}`) }) c.rate.Store(&ratelimit.Controller{InitialBackoff: time.Minute}) results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d) var limit *ratelimit.RateLimitError for _, result := range results { if result.Err == nil || !errors.As(result.Err, &limit) || strings.Contains(result.Err.Error(), "private") { t.Fatalf("row not failed as rate limit: %v", result.Err) } } again := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d) if again[0].Err == nil || !errors.As(again[0].Err, &limit) || calls != 1 { t.Fatalf("cooldown not armed: %v after %d calls", again[0].Err, calls) } } // A provider that rejects large schemas outright (Gemini's complexity cap // scales with the registry) must not fail the rows: the chunk halves until // accepted and the client remembers the working size. func TestBatchSplitsOnProviderSchemaRejection(t *testing.T) { f1, f2, d := batchRows() f3 := f1 f3.ID, f3.Fingerprint, f3.Counterparty = "tx_three", "fp_three", "Aral" f4 := f1 f4.ID, f4.Fingerprint, f4.Counterparty = "tx_four", "fp_four", "ALDI" calls, oversized := 0, 0 c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { calls++ prompt := decodeClassificationPrompt(t, r) category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food"))) if len(prompt.Transactions) > 2 { oversized++ w.WriteHeader(400) return } answers := make([]string, 0, len(prompt.Transactions)) for _, row := range prompt.Transactions { answers = append(answers, `{"ref":"`+row.Ref+`","merchant_id":null,"new_merchant":null,"category_id":"`+category+`","tag_ids":[],"confidence":"high"}`) } reply(w, `{"transactions":[`+strings.Join(answers, ",")+`]}`) }) results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2, f3, f4}, d) for i, result := range results { if result.Err != nil || result.Proposal.Enrichment.CategoryID != "cat_food" { t.Fatalf("row %d lost to schema rejection: %+v", i, result) } } if oversized != 1 || calls != 3 { t.Fatalf("expected one rejected probe then two halves, got %d calls (%d oversized)", calls, oversized) } if c.batchCap() != 2 { t.Fatalf("working batch size not learned: %d", c.batchCap()) } // The learned cap is respected up front on the next batch. before := calls _ = c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2, f3, f4}, d) if calls-before != 2 { t.Fatalf("learned cap ignored: %d extra calls", calls-before) } }