Use compact classification IDs and extend preview lifetime
This commit is contained in:
@@ -360,6 +360,65 @@ func TestPreviewIsReadOnlySelectedApplyPreservesFactsAndOtherFields(t *testing.T
|
|||||||
t.Fatal("consumed preview applied twice")
|
t.Fatal("consumed preview applied twice")
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func TestPreviewExpiresAfterTwentyFourHours(t *testing.T) {
|
||||||
|
for _, tc := range []struct {
|
||||||
|
name string
|
||||||
|
age time.Duration
|
||||||
|
newPreview bool
|
||||||
|
expired bool
|
||||||
|
}{
|
||||||
|
{name: "apply before expiry", age: 24*time.Hour - time.Minute},
|
||||||
|
{name: "apply after expiry", age: 24*time.Hour + time.Minute, expired: true},
|
||||||
|
{name: "new preview retains unexpired review", age: 24*time.Hour - time.Minute, newPreview: true},
|
||||||
|
{name: "new preview discards expired review", age: 24*time.Hour + time.Minute, newPreview: true, expired: true},
|
||||||
|
} {
|
||||||
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
|
ctx := context.Background()
|
||||||
|
a, s := testApp(t)
|
||||||
|
s = seed(t, a, s)
|
||||||
|
mockClassifier(t, a)
|
||||||
|
p, err := runPreview(t, a, PreviewRequest{Revision: s.Revision, From: "2026-09-01", To: "2026-09-30", Model: "test/model", Fields: Fields{Category: true}})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(p.Changes) != 2 {
|
||||||
|
t.Fatalf("expected two proposed changes: %+v", p)
|
||||||
|
}
|
||||||
|
a.mu.Lock()
|
||||||
|
p.created = time.Now().Add(-tc.age)
|
||||||
|
a.previews[p.ID] = p
|
||||||
|
a.mu.Unlock()
|
||||||
|
if tc.newPreview {
|
||||||
|
// Completing another run performs expired-preview cleanup.
|
||||||
|
// An empty range needs no additional provider request.
|
||||||
|
if _, err := runPreview(t, a, PreviewRequest{Revision: s.Revision, From: "2025-01-01", To: "2025-01-31", Model: "test/model", Fields: Fields{Category: true}}); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
change := p.Changes[0]
|
||||||
|
_, err = a.ApplyPreview(ctx, p.ID, p.Revision, []string{change.ID}, nil)
|
||||||
|
if (err != nil) != tc.expired {
|
||||||
|
t.Fatalf("apply at age %s: error = %v, expired = %t", tc.age, err, tc.expired)
|
||||||
|
}
|
||||||
|
after, err := a.Snapshot(ctx)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
expected := domain.Clone(s.Data)
|
||||||
|
if !tc.expired {
|
||||||
|
for i := range expected.Transactions {
|
||||||
|
if expected.Transactions[i].Facts.ID == change.ID {
|
||||||
|
expected.Transactions[i].Enrichment = change.After
|
||||||
|
}
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if !reflect.DeepEqual(after.Data, expected) {
|
||||||
|
t.Fatal("expiry handling did not preserve the expected transaction state")
|
||||||
|
}
|
||||||
|
})
|
||||||
|
}
|
||||||
|
}
|
||||||
func TestStalePreviewCannotOverwriteManualCorrection(t *testing.T) {
|
func TestStalePreviewCannotOverwriteManualCorrection(t *testing.T) {
|
||||||
a, s := testApp(t)
|
a, s := testApp(t)
|
||||||
s = seed(t, a, s)
|
s = seed(t, a, s)
|
||||||
@@ -491,7 +550,7 @@ func TestApplyPreviewHonoursReviewerEdits(t *testing.T) {
|
|||||||
func TestImportNeverAutoAppliesLowConfidenceCategory(t *testing.T) {
|
func TestImportNeverAutoAppliesLowConfidenceCategory(t *testing.T) {
|
||||||
a, s := testApp(t)
|
a, s := testApp(t)
|
||||||
provider := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
provider := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||||
content := `{"merchant_id":null,"new_merchant":"REWE","category_id":"groceries","tag_ids":[],"confidence":"low"}`
|
content := `{"merchant_id":null,"new_merchant":"REWE","category_id":"c1","tag_ids":[],"confidence":"low"}`
|
||||||
json.NewEncoder(w).Encode(map[string]any{"choices": []any{map[string]any{
|
json.NewEncoder(w).Encode(map[string]any{"choices": []any{map[string]any{
|
||||||
"finish_reason": "stop",
|
"finish_reason": "stop",
|
||||||
"message": map[string]any{"content": content},
|
"message": map[string]any{"content": content},
|
||||||
|
|||||||
@@ -13,6 +13,8 @@ import (
|
|||||||
"finance-duck/internal/domain"
|
"finance-duck/internal/domain"
|
||||||
)
|
)
|
||||||
|
|
||||||
|
const previewLifetime = 24 * time.Hour
|
||||||
|
|
||||||
type Fields struct {
|
type Fields struct {
|
||||||
Merchant bool `json:"merchant"`
|
Merchant bool `json:"merchant"`
|
||||||
Category bool `json:"category"`
|
Category bool `json:"category"`
|
||||||
@@ -173,7 +175,7 @@ func (a *App) runPreview(ctx context.Context, cancel context.CancelFunc, client
|
|||||||
return
|
return
|
||||||
}
|
}
|
||||||
for id, old := range a.previews {
|
for id, old := range a.previews {
|
||||||
if time.Since(old.created) > time.Hour {
|
if time.Since(old.created) > previewLifetime {
|
||||||
delete(a.previews, id)
|
delete(a.previews, id)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -343,7 +345,7 @@ func (a *App) ApplyPreview(ctx context.Context, id, rev string, ids []string, ed
|
|||||||
a.mu.Lock()
|
a.mu.Lock()
|
||||||
defer a.mu.Unlock()
|
defer a.mu.Unlock()
|
||||||
p, ok := a.previews[id]
|
p, ok := a.previews[id]
|
||||||
if !ok || time.Since(p.created) > time.Hour {
|
if !ok || time.Since(p.created) > previewLifetime {
|
||||||
return State{}, errors.New("preview expired or unknown; analyse again")
|
return State{}, errors.New("preview expired or unknown; analyse again")
|
||||||
}
|
}
|
||||||
if rev != p.Revision {
|
if rev != p.Revision {
|
||||||
|
|||||||
@@ -129,7 +129,7 @@ func (c *Client) ClassifyBatch(ctx context.Context, rows []domain.Facts, data do
|
|||||||
payload.Transactions = append(payload.Transactions, row)
|
payload.Transactions = append(payload.Transactions, row)
|
||||||
similar.WriteString(f.RawDescription + " " + f.Counterparty + " ")
|
similar.WriteString(f.RawDescription + " " + f.Counterparty + " ")
|
||||||
}
|
}
|
||||||
payload.History = history(domain.Facts{RawDescription: similar.String()}, data, clean, 40)
|
payload.History = candidates.history(domain.Facts{RawDescription: similar.String()}, data, clean, 40)
|
||||||
payload.Categories = candidates.categories
|
payload.Categories = candidates.categories
|
||||||
payload.Tags = candidates.tags
|
payload.Tags = candidates.tags
|
||||||
payload.Merchants = candidates.merchants
|
payload.Merchants = candidates.merchants
|
||||||
|
|||||||
@@ -2,7 +2,6 @@ package classification
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
"encoding/json"
|
|
||||||
"errors"
|
"errors"
|
||||||
"io"
|
"io"
|
||||||
"net/http"
|
"net/http"
|
||||||
@@ -32,28 +31,34 @@ func TestBatchClassifiesEveryRowInOneRequest(t *testing.T) {
|
|||||||
calls := 0
|
calls := 0
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
calls++
|
calls++
|
||||||
var req struct {
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
Messages []struct {
|
if len(prompt.Transactions) != 2 {
|
||||||
Content string `json:"content"`
|
t.Errorf("batch prompt missing transactions: %+v", prompt.Transactions)
|
||||||
} `json:"messages"`
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
}
|
|
||||||
if json.NewDecoder(r.Body).Decode(&req) != nil || len(req.Messages) != 2 {
|
|
||||||
w.WriteHeader(400)
|
|
||||||
return
|
return
|
||||||
}
|
}
|
||||||
var prompt struct {
|
category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food")))
|
||||||
Transactions []struct{ Ref, Counterparty, Amount, Currency string } `json:"transactions"`
|
merchant, tag := "", ""
|
||||||
|
for _, candidate := range prompt.Merchants {
|
||||||
|
if candidate.Name == "coffee house" {
|
||||||
|
merchant = candidate.ID
|
||||||
}
|
}
|
||||||
if json.Unmarshal([]byte(req.Messages[1].Content), &prompt) != nil || len(prompt.Transactions) != 2 {
|
}
|
||||||
t.Errorf("batch prompt missing transactions: %s", req.Messages[1].Content)
|
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 {
|
for _, row := range prompt.Transactions {
|
||||||
if row.Amount == "" || row.Currency != "EUR" {
|
if row.Amount == "" || row.Currency != "EUR" {
|
||||||
t.Errorf("row %s lost amount or currency: %+v", row.Ref, row)
|
t.Errorf("row %s lost amount or currency: %+v", row.Ref, row)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
reply(w, `{"transactions":[{"ref":"r1","merchant_id":"mer_coffee","new_merchant":null,"category_id":"cat_food","tag_ids":["tag_daily"],"confidence":"high"},`+
|
reply(w, `{"transactions":[{"ref":"`+prompt.Transactions[1].Ref+`","merchant_id":null,"new_merchant":"Kleins Backstube","category_id":"`+category+`","tag_ids":[],"confidence":"medium"},`+
|
||||||
`{"ref":"r2","merchant_id":null,"new_merchant":"Kleins Backstube","category_id":"cat_food","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)
|
results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d)
|
||||||
if calls != 1 {
|
if calls != 1 {
|
||||||
@@ -71,6 +76,8 @@ func TestBatchClassifiesEveryRowInOneRequest(t *testing.T) {
|
|||||||
if second.NewMerchant == nil || second.NewMerchant.Name != "Kleins Backstube" ||
|
if second.NewMerchant == nil || second.NewMerchant.Name != "Kleins Backstube" ||
|
||||||
!reflect.DeepEqual(second.NewMerchant.Aliases, []string{"Kleins Backstube"}) ||
|
!reflect.DeepEqual(second.NewMerchant.Aliases, []string{"Kleins Backstube"}) ||
|
||||||
second.Enrichment.MerchantID != second.NewMerchant.ID ||
|
second.Enrichment.MerchantID != second.NewMerchant.ID ||
|
||||||
|
second.Enrichment.CategoryID != "cat_food" ||
|
||||||
|
len(second.Enrichment.TagIDs) != 0 ||
|
||||||
second.Enrichment.Classification.Confidence != "medium" {
|
second.Enrichment.Classification.Confidence != "medium" {
|
||||||
t.Fatalf("second row lost: %+v", second)
|
t.Fatalf("second row lost: %+v", second)
|
||||||
}
|
}
|
||||||
@@ -80,8 +87,10 @@ func TestBatchClassifiesEveryRowInOneRequest(t *testing.T) {
|
|||||||
func TestBatchIsolatesInvalidRows(t *testing.T) {
|
func TestBatchIsolatesInvalidRows(t *testing.T) {
|
||||||
f1, f2, d := batchRows()
|
f1, f2, d := batchRows()
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
reply(w, `{"transactions":[{"ref":"r1","merchant_id":null,"new_merchant":null,"category_id":"cat_food","tag_ids":[],"confidence":"high"},`+
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
`{"ref":"r2","merchant_id":null,"new_merchant":null,"category_id":"cat_forged","tag_ids":[],"confidence":"high"}]}`)
|
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)
|
results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d)
|
||||||
if results[0].Err != nil || results[0].Proposal.Enrichment.CategoryID != "cat_food" {
|
if results[0].Err != nil || results[0].Proposal.Enrichment.CategoryID != "cat_food" {
|
||||||
@@ -96,8 +105,10 @@ func TestBatchIsolatesInvalidRows(t *testing.T) {
|
|||||||
func TestBatchSharesOneMintedMerchant(t *testing.T) {
|
func TestBatchSharesOneMintedMerchant(t *testing.T) {
|
||||||
f1, f2, d := batchRows()
|
f1, f2, d := batchRows()
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
reply(w, `{"transactions":[{"ref":"r1","merchant_id":null,"new_merchant":"REWE","category_id":"cat_food","tag_ids":[],"confidence":"high"},`+
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
`{"ref":"r2","merchant_id":null,"new_merchant":"REWE","category_id":"cat_food","tag_ids":[],"confidence":"high"}]}`)
|
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)
|
results := c.ClassifyBatch(context.Background(), []domain.Facts{f1, f2}, d)
|
||||||
if results[0].Err != nil || results[1].Err != nil {
|
if results[0].Err != nil || results[1].Err != nil {
|
||||||
@@ -144,19 +155,8 @@ func TestBatchSplitsOnProviderSchemaRejection(t *testing.T) {
|
|||||||
calls, oversized := 0, 0
|
calls, oversized := 0, 0
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
calls++
|
calls++
|
||||||
var req struct {
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
Messages []struct {
|
category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food")))
|
||||||
Content string `json:"content"`
|
|
||||||
} `json:"messages"`
|
|
||||||
}
|
|
||||||
if json.NewDecoder(r.Body).Decode(&req) != nil {
|
|
||||||
w.WriteHeader(500)
|
|
||||||
return
|
|
||||||
}
|
|
||||||
var prompt struct {
|
|
||||||
Transactions []struct{ Ref string } `json:"transactions"`
|
|
||||||
}
|
|
||||||
_ = json.Unmarshal([]byte(req.Messages[1].Content), &prompt)
|
|
||||||
if len(prompt.Transactions) > 2 {
|
if len(prompt.Transactions) > 2 {
|
||||||
oversized++
|
oversized++
|
||||||
w.WriteHeader(400)
|
w.WriteHeader(400)
|
||||||
@@ -164,7 +164,7 @@ func TestBatchSplitsOnProviderSchemaRejection(t *testing.T) {
|
|||||||
}
|
}
|
||||||
answers := make([]string, 0, len(prompt.Transactions))
|
answers := make([]string, 0, len(prompt.Transactions))
|
||||||
for _, row := range prompt.Transactions {
|
for _, row := range prompt.Transactions {
|
||||||
answers = append(answers, `{"ref":"`+row.Ref+`","merchant_id":null,"new_merchant":null,"category_id":"cat_food","tag_ids":[],"confidence":"high"}`)
|
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, ",")+`]}`)
|
reply(w, `{"transactions":[`+strings.Join(answers, ",")+`]}`)
|
||||||
})
|
})
|
||||||
|
|||||||
@@ -3,6 +3,7 @@ package classification
|
|||||||
import (
|
import (
|
||||||
"slices"
|
"slices"
|
||||||
"sort"
|
"sort"
|
||||||
|
"strconv"
|
||||||
"strings"
|
"strings"
|
||||||
"unicode"
|
"unicode"
|
||||||
|
|
||||||
@@ -150,8 +151,6 @@ type merchantPrompt struct {
|
|||||||
UsualCategory string `json:"usual_category,omitempty"`
|
UsualCategory string `json:"usual_category,omitempty"`
|
||||||
}
|
}
|
||||||
|
|
||||||
// candidate is the historical merchant prompt shape used by older callers.
|
|
||||||
type candidate = merchantPrompt
|
|
||||||
type promptHistory struct {
|
type promptHistory struct {
|
||||||
Date string `json:"date"`
|
Date string `json:"date"`
|
||||||
Amount string `json:"amount"`
|
Amount string `json:"amount"`
|
||||||
@@ -172,6 +171,9 @@ type candidateSet struct {
|
|||||||
categoryIDs map[string]string
|
categoryIDs map[string]string
|
||||||
tagIDs map[string]string
|
tagIDs map[string]string
|
||||||
merchantIDs map[string]string
|
merchantIDs map[string]string
|
||||||
|
categoryRefs map[string]string
|
||||||
|
tagRefs map[string]string
|
||||||
|
merchantRefs map[string]string
|
||||||
}
|
}
|
||||||
|
|
||||||
func similarity(description, name string) int {
|
func similarity(description, name string) int {
|
||||||
@@ -194,9 +196,8 @@ func similarity(description, name string) int {
|
|||||||
return score
|
return score
|
||||||
}
|
}
|
||||||
|
|
||||||
// retrieve emits every registry entry with its real id. The legacy cleaner
|
// retrieve offers every eligible registry entry under a short request-local
|
||||||
// arguments remain in the signature because CSV/classification fixtures use
|
// reference. Names and paths retain their meaning; canonical IDs stay local.
|
||||||
// this helper directly; ranking and bounding are intentionally gone.
|
|
||||||
func retrieve(_ string, kind string, data domain.Dataset, clean, merchantClean func(string) string) candidateSet {
|
func retrieve(_ string, kind string, data domain.Dataset, clean, merchantClean func(string) string) candidateSet {
|
||||||
parents := map[string]bool{}
|
parents := map[string]bool{}
|
||||||
for _, cat := range data.Categories {
|
for _, cat := range data.Categories {
|
||||||
@@ -206,6 +207,9 @@ func retrieve(_ string, kind string, data domain.Dataset, clean, merchantClean f
|
|||||||
categoryIDs: map[string]string{},
|
categoryIDs: map[string]string{},
|
||||||
tagIDs: map[string]string{},
|
tagIDs: map[string]string{},
|
||||||
merchantIDs: map[string]string{},
|
merchantIDs: map[string]string{},
|
||||||
|
categoryRefs: map[string]string{},
|
||||||
|
tagRefs: map[string]string{},
|
||||||
|
merchantRefs: map[string]string{},
|
||||||
}
|
}
|
||||||
for _, cat := range data.Categories {
|
for _, cat := range data.Categories {
|
||||||
if cat.Kind != kind || parents[cat.ID] {
|
if cat.Kind != kind || parents[cat.ID] {
|
||||||
@@ -216,19 +220,31 @@ func retrieve(_ string, kind string, data domain.Dataset, clean, merchantClean f
|
|||||||
path = clean(path)
|
path = clean(path)
|
||||||
}
|
}
|
||||||
set.categories = append(set.categories, categoryPrompt{ID: cat.ID, Path: path, Kind: cat.Kind, Hint: cleanText(clean, cat.Hint)})
|
set.categories = append(set.categories, categoryPrompt{ID: cat.ID, Path: path, Kind: cat.Kind, Hint: cleanText(clean, cat.Hint)})
|
||||||
set.categoryIDs[cat.ID] = cat.ID
|
|
||||||
}
|
}
|
||||||
sort.Slice(set.categories, func(i, j int) bool {
|
sort.Slice(set.categories, func(i, j int) bool {
|
||||||
return set.categories[i].Path < set.categories[j].Path || set.categories[i].Path == set.categories[j].Path && set.categories[i].ID < set.categories[j].ID
|
return set.categories[i].Path < set.categories[j].Path || set.categories[i].Path == set.categories[j].Path && set.categories[i].ID < set.categories[j].ID
|
||||||
})
|
})
|
||||||
|
for i := range set.categories {
|
||||||
|
category := &set.categories[i]
|
||||||
|
ref := "c" + strconv.Itoa(i+1)
|
||||||
|
set.categoryIDs[ref] = category.ID
|
||||||
|
set.categoryRefs[category.ID] = ref
|
||||||
|
category.ID = ref
|
||||||
|
}
|
||||||
for _, tag := range data.Tags {
|
for _, tag := range data.Tags {
|
||||||
name := cleanText(clean, tag.Name)
|
name := cleanText(clean, tag.Name)
|
||||||
set.tags = append(set.tags, tagPrompt{ID: tag.ID, Name: name, Hint: cleanText(clean, tag.Hint)})
|
set.tags = append(set.tags, tagPrompt{ID: tag.ID, Name: name, Hint: cleanText(clean, tag.Hint)})
|
||||||
set.tagIDs[tag.ID] = tag.ID
|
|
||||||
}
|
}
|
||||||
sort.Slice(set.tags, func(i, j int) bool {
|
sort.Slice(set.tags, func(i, j int) bool {
|
||||||
return set.tags[i].Name < set.tags[j].Name || set.tags[i].Name == set.tags[j].Name && set.tags[i].ID < set.tags[j].ID
|
return set.tags[i].Name < set.tags[j].Name || set.tags[i].Name == set.tags[j].Name && set.tags[i].ID < set.tags[j].ID
|
||||||
})
|
})
|
||||||
|
for i := range set.tags {
|
||||||
|
tag := &set.tags[i]
|
||||||
|
ref := "t" + strconv.Itoa(i+1)
|
||||||
|
set.tagIDs[ref] = tag.ID
|
||||||
|
set.tagRefs[tag.ID] = ref
|
||||||
|
tag.ID = ref
|
||||||
|
}
|
||||||
usual := map[string]string{}
|
usual := map[string]string{}
|
||||||
counts := map[string]map[string]int{}
|
counts := map[string]map[string]int{}
|
||||||
for _, tx := range data.Transactions {
|
for _, tx := range data.Transactions {
|
||||||
@@ -263,13 +279,19 @@ func retrieve(_ string, kind string, data domain.Dataset, clean, merchantClean f
|
|||||||
}
|
}
|
||||||
set.merchants = append(set.merchants, merchantPrompt{
|
set.merchants = append(set.merchants, merchantPrompt{
|
||||||
ID: merchant.ID, Name: name, Aliases: aliases,
|
ID: merchant.ID, Name: name, Aliases: aliases,
|
||||||
UsualCategory: usualCategory,
|
UsualCategory: set.categoryRefs[usualCategory],
|
||||||
})
|
})
|
||||||
set.merchantIDs[merchant.ID] = merchant.ID
|
|
||||||
}
|
}
|
||||||
sort.Slice(set.merchants, func(i, j int) bool {
|
sort.Slice(set.merchants, func(i, j int) bool {
|
||||||
return set.merchants[i].Name < set.merchants[j].Name || set.merchants[i].Name == set.merchants[j].Name && set.merchants[i].ID < set.merchants[j].ID
|
return set.merchants[i].Name < set.merchants[j].Name || set.merchants[i].Name == set.merchants[j].Name && set.merchants[i].ID < set.merchants[j].ID
|
||||||
})
|
})
|
||||||
|
for i := range set.merchants {
|
||||||
|
merchant := &set.merchants[i]
|
||||||
|
ref := "m" + strconv.Itoa(i+1)
|
||||||
|
set.merchantIDs[ref] = merchant.ID
|
||||||
|
set.merchantRefs[merchant.ID] = ref
|
||||||
|
merchant.ID = ref
|
||||||
|
}
|
||||||
return set
|
return set
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -314,16 +336,12 @@ func candidateIDs(values []categoryPrompt) []string {
|
|||||||
return ids
|
return ids
|
||||||
}
|
}
|
||||||
|
|
||||||
func answerSchema(d domain.Dataset, kind string) map[string]any {
|
// history selects precedent whose category is offered in this request: the
|
||||||
return retrieve("", kind, d, nil, nil).schema()
|
// nearest rows by word overlap, filled out with the most recent. The user's
|
||||||
}
|
// own decisions — manual edits and locally applied merchant rules — outrank
|
||||||
|
// rows the model classified itself. References use the same mapping as the
|
||||||
// history selects precedent for the prompt: the nearest rows by word overlap,
|
// candidate lists and response schema.
|
||||||
// filled out with the most recent. The user's own decisions — manual edits
|
func (c candidateSet) history(f domain.Facts, d domain.Dataset, clean func(string) string, limit int) []promptHistory {
|
||||||
// and locally applied merchant rules — outrank rows the model classified
|
|
||||||
// itself, so one correction beats any number of uncorrected AI answers for
|
|
||||||
// the same payee.
|
|
||||||
func history(f domain.Facts, d domain.Dataset, clean func(string) string, limit int) []promptHistory {
|
|
||||||
type row struct {
|
type row struct {
|
||||||
tx domain.Transaction
|
tx domain.Transaction
|
||||||
score int
|
score int
|
||||||
@@ -332,7 +350,7 @@ func history(f domain.Facts, d domain.Dataset, clean func(string) string, limit
|
|||||||
rows := []row{}
|
rows := []row{}
|
||||||
for _, tx := range d.Transactions {
|
for _, tx := range d.Transactions {
|
||||||
e := tx.Enrichment
|
e := tx.Enrichment
|
||||||
if tx.Facts.ID == f.ID || e.Kind == "transfer" || e.CategoryID == "" || e.CategoryID == domain.ExpenseFallback || e.CategoryID == domain.IncomeFallback {
|
if tx.Facts.ID == f.ID || e.Kind == "transfer" || c.categoryRefs[e.CategoryID] == "" || e.CategoryID == domain.ExpenseFallback || e.CategoryID == domain.IncomeFallback {
|
||||||
continue
|
continue
|
||||||
}
|
}
|
||||||
source := tx.Enrichment.Classification.Source
|
source := tx.Enrichment.Classification.Source
|
||||||
@@ -382,9 +400,11 @@ func history(f domain.Facts, d domain.Dataset, clean func(string) string, limit
|
|||||||
}
|
}
|
||||||
out := make([]promptHistory, 0, len(rows))
|
out := make([]promptHistory, 0, len(rows))
|
||||||
for _, row := range rows {
|
for _, row := range rows {
|
||||||
tags := row.tx.Enrichment.TagIDs
|
tags := make([]string, 0, len(row.tx.Enrichment.TagIDs))
|
||||||
if tags == nil {
|
for _, id := range row.tx.Enrichment.TagIDs {
|
||||||
tags = []string{}
|
if ref := c.tagRefs[id]; ref != "" {
|
||||||
|
tags = append(tags, ref)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
source := "ai"
|
source := "ai"
|
||||||
if row.user {
|
if row.user {
|
||||||
@@ -393,8 +413,8 @@ func history(f domain.Facts, d domain.Dataset, clean func(string) string, limit
|
|||||||
out = append(out, promptHistory{
|
out = append(out, promptHistory{
|
||||||
Date: row.tx.Facts.BookingDate, Amount: string(row.tx.Facts.Amount),
|
Date: row.tx.Facts.BookingDate, Amount: string(row.tx.Facts.Amount),
|
||||||
Description: clean(row.tx.Facts.RawDescription), Counterparty: clean(row.tx.Facts.Counterparty),
|
Description: clean(row.tx.Facts.RawDescription), Counterparty: clean(row.tx.Facts.Counterparty),
|
||||||
CategoryID: row.tx.Enrichment.CategoryID, MerchantID: row.tx.Enrichment.MerchantID,
|
CategoryID: c.categoryRefs[row.tx.Enrichment.CategoryID], MerchantID: c.merchantRefs[row.tx.Enrichment.MerchantID],
|
||||||
TagIDs: append([]string{}, tags...),
|
TagIDs: tags,
|
||||||
Source: source,
|
Source: source,
|
||||||
})
|
})
|
||||||
}
|
}
|
||||||
|
|||||||
@@ -220,7 +220,7 @@ func (c *Client) Classify(ctx context.Context, facts domain.Facts, data domain.D
|
|||||||
userPayload.Transaction.Counterparty = clean(facts.Counterparty)
|
userPayload.Transaction.Counterparty = clean(facts.Counterparty)
|
||||||
userPayload.Transaction.Account.Institution = clean(institution)
|
userPayload.Transaction.Account.Institution = clean(institution)
|
||||||
userPayload.Transaction.Account.Currency = facts.Currency
|
userPayload.Transaction.Account.Currency = facts.Currency
|
||||||
userPayload.History = history(facts, data, clean, 40)
|
userPayload.History = candidates.history(facts, data, clean, 40)
|
||||||
userPayload.Categories = candidates.categories
|
userPayload.Categories = candidates.categories
|
||||||
userPayload.Tags = candidates.tags
|
userPayload.Tags = candidates.tags
|
||||||
userPayload.Merchants = candidates.merchants
|
userPayload.Merchants = candidates.merchants
|
||||||
|
|||||||
@@ -28,7 +28,7 @@ func fixture() (domain.Facts, domain.Dataset) {
|
|||||||
return f, d
|
return f, d
|
||||||
}
|
}
|
||||||
|
|
||||||
const validAnswer = `{"merchant_id":null,"new_merchant":null,"category_id":"cat_food","tag_ids":[],"confidence":"medium"}`
|
const validAnswer = `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"medium"}`
|
||||||
|
|
||||||
func reply(w http.ResponseWriter, content string) {
|
func reply(w http.ResponseWriter, content string) {
|
||||||
w.Header().Set("Content-Type", "application/json")
|
w.Header().Set("Content-Type", "application/json")
|
||||||
@@ -44,6 +44,39 @@ func mockClient(t *testing.T, handler http.HandlerFunc) *Client {
|
|||||||
return client
|
return client
|
||||||
}
|
}
|
||||||
|
|
||||||
|
type classificationPrompt struct {
|
||||||
|
Categories []categoryPrompt `json:"categories"`
|
||||||
|
Tags []tagPrompt `json:"tags"`
|
||||||
|
Merchants []merchantPrompt `json:"merchants"`
|
||||||
|
History []promptHistory `json:"history"`
|
||||||
|
Transactions []struct {
|
||||||
|
Ref string `json:"ref"`
|
||||||
|
Counterparty string `json:"counterparty"`
|
||||||
|
Amount string `json:"amount"`
|
||||||
|
Currency string `json:"currency"`
|
||||||
|
} `json:"transactions"`
|
||||||
|
}
|
||||||
|
|
||||||
|
func decodeClassificationPrompt(t *testing.T, r *http.Request) classificationPrompt {
|
||||||
|
t.Helper()
|
||||||
|
var req struct {
|
||||||
|
Messages []struct {
|
||||||
|
Content string `json:"content"`
|
||||||
|
} `json:"messages"`
|
||||||
|
}
|
||||||
|
if err := json.NewDecoder(r.Body).Decode(&req); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if len(req.Messages) != 2 {
|
||||||
|
t.Fatalf("expected system and user messages, got %d", len(req.Messages))
|
||||||
|
}
|
||||||
|
var prompt classificationPrompt
|
||||||
|
if err := json.Unmarshal([]byte(req.Messages[1].Content), &prompt); err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
return prompt
|
||||||
|
}
|
||||||
|
|
||||||
func TestExplicitDefaultsAreOptInAndBypassAI(t *testing.T) {
|
func TestExplicitDefaultsAreOptInAndBypassAI(t *testing.T) {
|
||||||
f, d := fixture()
|
f, d := fixture()
|
||||||
d.Merchants[0].UseDefaults = true
|
d.Merchants[0].UseDefaults = true
|
||||||
@@ -72,7 +105,22 @@ func TestForceAIOverridesRuleWithoutChangingKind(t *testing.T) {
|
|||||||
f, d := fixture()
|
f, d := fixture()
|
||||||
d.Merchants[0].UseDefaults = true
|
d.Merchants[0].UseDefaults = true
|
||||||
calls := 0
|
calls := 0
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { calls++; reply(w, validAnswer) })
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
calls++
|
||||||
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
|
// Food is an expense-only choice; do not reuse c1 after the request
|
||||||
|
// switches to income, where that reference names a different category.
|
||||||
|
categoryID := "c999"
|
||||||
|
for _, category := range prompt.Categories {
|
||||||
|
if category.Path == normalize(domain.CategoryPath(d, "cat_food")) {
|
||||||
|
categoryID = category.ID
|
||||||
|
}
|
||||||
|
if calls == 2 && category.Kind != "income" {
|
||||||
|
t.Errorf("income request offered an expense category: %+v", category)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
reply(w, fmt.Sprintf(`{"merchant_id":null,"new_merchant":null,"category_id":%q,"tag_ids":[],"confidence":"medium"}`, categoryID))
|
||||||
|
})
|
||||||
p, err := c.Classify(context.Background(), f, d, true)
|
p, err := c.Classify(context.Background(), f, d, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -82,7 +130,7 @@ func TestForceAIOverridesRuleWithoutChangingKind(t *testing.T) {
|
|||||||
}
|
}
|
||||||
f.Amount = "918.27"
|
f.Amount = "918.27"
|
||||||
p, err = c.Classify(context.Background(), f, d, true)
|
p, err = c.Classify(context.Background(), f, d, true)
|
||||||
if err == nil || p.Enrichment.Kind != "income" || p.Enrichment.CategoryID != domain.IncomeFallback {
|
if err == nil || calls != 2 || p.Enrichment.Kind != "income" || p.Enrichment.CategoryID != domain.IncomeFallback {
|
||||||
t.Fatalf("income sign: %+v %v", p, err)
|
t.Fatalf("income sign: %+v %v", p, err)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
@@ -117,21 +165,24 @@ func TestTransferNeverCallsAIOrAliases(t *testing.T) {
|
|||||||
|
|
||||||
func TestInvalidModelOutputsFailClosed(t *testing.T) {
|
func TestInvalidModelOutputsFailClosed(t *testing.T) {
|
||||||
cases := map[string]string{
|
cases := map[string]string{
|
||||||
"unknown key": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":0.9}`,
|
"unknown key": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high","unexpected":true}`,
|
||||||
"change kind": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"kind":"transfer"}`,
|
"change kind": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high","kind":"transfer"}`,
|
||||||
"missing field": `{"merchant_id":null,"category_id":"c1","tag_ids":[]}`,
|
"missing field": `{"merchant_id":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
"duplicate key": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","category_id":"c2","tag_ids":[]}`,
|
"duplicate key": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","category_id":"c2","tag_ids":[],"confidence":"high"}`,
|
||||||
"case folded key": `{"Merchant_ID":null,"new_merchant":null,"category_id":"c1","tag_ids":[]}`,
|
"case folded key": `{"Merchant_ID":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
"unknown category": `{"merchant_id":null,"new_merchant":null,"category_id":"cat_invented","tag_ids":[]}`,
|
"unknown category": `{"merchant_id":null,"new_merchant":null,"category_id":"c999","tag_ids":[],"confidence":"high"}`,
|
||||||
"real ID not offered": `{"merchant_id":null,"new_merchant":null,"category_id":"cat_food","tag_ids":[]}`,
|
"canonical category": `{"merchant_id":null,"new_merchant":null,"category_id":"cat_food","tag_ids":[],"confidence":"high"}`,
|
||||||
"unknown tag": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":["t999"]}`,
|
"unknown tag": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":["t999"],"confidence":"high"}`,
|
||||||
"duplicate tags": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":["t1","t1"]}`,
|
"canonical tag": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":["tag_daily"],"confidence":"high"}`,
|
||||||
"null tags": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":null}`,
|
"duplicate tags": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":["t1","t1"],"confidence":"high"}`,
|
||||||
"null tag member": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[null]}`,
|
"null tags": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":null,"confidence":"high"}`,
|
||||||
"unknown merchant": `{"merchant_id":"m999","new_merchant":null,"category_id":"c1","tag_ids":[]}`,
|
"null tag member": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[null],"confidence":"high"}`,
|
||||||
"both merchant modes": `{"merchant_id":"m1","new_merchant":"Coffee","category_id":"c1","tag_ids":[]}`,
|
"unknown merchant": `{"merchant_id":"m999","new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
"blank proposal": `{"merchant_id":null,"new_merchant":" ","category_id":"c1","tag_ids":[]}`,
|
"canonical merchant": `{"merchant_id":"mer_coffee","new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
"wrong scalar": `{"merchant_id":23,"new_merchant":null,"category_id":"c1","tag_ids":[]}`,
|
"both merchant modes": `{"merchant_id":"m1","new_merchant":"Coffee","category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
|
"blank proposal": `{"merchant_id":null,"new_merchant":" ","category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
|
"wrong scalar": `{"merchant_id":23,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`,
|
||||||
|
"numeric confidence": `{"merchant_id":null,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":0.9}`,
|
||||||
"trailing JSON": validAnswer + ` {}`,
|
"trailing JSON": validAnswer + ` {}`,
|
||||||
"markdown": "```json\n" + validAnswer + "\n```",
|
"markdown": "```json\n" + validAnswer + "\n```",
|
||||||
"array": "[" + validAnswer + "]",
|
"array": "[" + validAnswer + "]",
|
||||||
@@ -157,9 +208,9 @@ func TestMerchantSelectionAndLocalProposal(t *testing.T) {
|
|||||||
name, content, merchant string
|
name, content, merchant string
|
||||||
new bool
|
new bool
|
||||||
}{
|
}{
|
||||||
{"existing", `{"merchant_id":"mer_coffee","new_merchant":null,"category_id":"cat_food","tag_ids":["tag_daily"],"confidence":"high"}`, "mer_coffee", false},
|
{"existing", `{"merchant_id":"m1","new_merchant":null,"category_id":"c1","tag_ids":["t1"],"confidence":"high"}`, "mer_coffee", false},
|
||||||
{"duplicate alias", `{"merchant_id":null,"new_merchant":"COFFEE-house","category_id":"cat_food","tag_ids":["tag_daily"],"confidence":"high"}`, "mer_coffee", false},
|
{"duplicate alias", `{"merchant_id":null,"new_merchant":"COFFEE-house","category_id":"c1","tag_ids":["t1"],"confidence":"high"}`, "mer_coffee", false},
|
||||||
{"new", `{"merchant_id":null,"new_merchant":"Bakery Lane","category_id":"cat_food","tag_ids":["tag_daily"],"confidence":"high"}`, "", true},
|
{"new", `{"merchant_id":null,"new_merchant":"Bakery Lane","category_id":"c1","tag_ids":["t1"],"confidence":"high"}`, "", true},
|
||||||
}
|
}
|
||||||
for _, tc := range cases {
|
for _, tc := range cases {
|
||||||
t.Run(tc.name, func(t *testing.T) {
|
t.Run(tc.name, func(t *testing.T) {
|
||||||
@@ -216,18 +267,14 @@ func TestIdentifierOnlyPromptRedactionAndRouting(t *testing.T) {
|
|||||||
if len(messages) != 2 {
|
if len(messages) != 2 {
|
||||||
t.Fatal("unexpected messages")
|
t.Fatal("unexpected messages")
|
||||||
}
|
}
|
||||||
var prompt struct {
|
wire, err := json.Marshal(captured)
|
||||||
Transaction map[string]any `json:"transaction"`
|
if err != nil {
|
||||||
History []any `json:"history"`
|
|
||||||
Categories []any `json:"categories"`
|
|
||||||
Tags []any `json:"tags"`
|
|
||||||
Merchants []any `json:"merchants"`
|
|
||||||
}
|
|
||||||
if err := json.Unmarshal([]byte(messages[1].Content), &prompt); err != nil {
|
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
}
|
}
|
||||||
if len(prompt.Transaction) == 0 || len(prompt.Categories) == 0 || len(prompt.Merchants) == 0 {
|
for _, canonicalID := range []string{"cat_food", "cat_expenses", "cat_income", "mer_coffee", "tag_daily"} {
|
||||||
t.Fatal("complete structured prompt missing")
|
if strings.Contains(string(wire), canonicalID) {
|
||||||
|
t.Errorf("request or response schema exposed canonical ID %q", canonicalID)
|
||||||
|
}
|
||||||
}
|
}
|
||||||
lower := strings.ToLower(messages[1].Content)
|
lower := strings.ToLower(messages[1].Content)
|
||||||
for _, secret := range []string{"private_external", "private_fingerprint", "tx_private", "account_private", "ext_local_secret", "private_source", "personal checking", "550e8400", "cobadeff", "secretpayment", "example.com", "alice privateperson", "de89370400440532013000", "de44500105175407324931"} {
|
for _, secret := range []string{"private_external", "private_fingerprint", "tx_private", "account_private", "ext_local_secret", "private_source", "personal checking", "550e8400", "cobadeff", "secretpayment", "example.com", "alice privateperson", "de89370400440532013000", "de44500105175407324931"} {
|
||||||
@@ -300,7 +347,7 @@ func TestUnsafeMerchantProposalDroppedWithoutLosingClassification(t *testing.T)
|
|||||||
t.Run(name, func(t *testing.T) {
|
t.Run(name, func(t *testing.T) {
|
||||||
f, d := fixture()
|
f, d := fixture()
|
||||||
f.Counterparty = "Alice Privateperson"
|
f.Counterparty = "Alice Privateperson"
|
||||||
answer, _ := json.Marshal(map[string]any{"merchant_id": nil, "new_merchant": name, "category_id": "cat_food", "tag_ids": []string{}, "confidence": "high"})
|
answer, _ := json.Marshal(map[string]any{"merchant_id": nil, "new_merchant": name, "category_id": "c1", "tag_ids": []string{}, "confidence": "high"})
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { reply(w, string(answer)) })
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) { reply(w, string(answer)) })
|
||||||
c.PrivateNames = []string{"Alice Privateperson"}
|
c.PrivateNames = []string{"Alice Privateperson"}
|
||||||
p, err := c.Classify(context.Background(), f, d, true)
|
p, err := c.Classify(context.Background(), f, d, true)
|
||||||
@@ -410,33 +457,67 @@ func TestCompleteRegistryPayloadAndGlobalDuplicateDetection(t *testing.T) {
|
|||||||
d.Tags = append(d.Tags, domain.Tag{ID: fmt.Sprintf("tag_%02d", i), Name: fmt.Sprintf("Tag %02d", i)})
|
d.Tags = append(d.Tags, domain.Tag{ID: fmt.Sprintf("tag_%02d", i), Name: fmt.Sprintf("Tag %02d", i)})
|
||||||
d.Categories = append(d.Categories, domain.Category{ID: fmt.Sprintf("cat_%02d", i), Name: fmt.Sprintf("Category %02d", i), Kind: "expense", ParentID: "cat_expenses"})
|
d.Categories = append(d.Categories, domain.Category{ID: fmt.Sprintf("cat_%02d", i), Name: fmt.Sprintf("Category %02d", i), Kind: "expense", ParentID: "cat_expenses"})
|
||||||
}
|
}
|
||||||
d.Merchants[34].Name = "Distant Bakery"
|
d.Merchants[34].Name = "Z Distant Bakery"
|
||||||
set := retrieve(f.RawDescription, "expense", d, redactor(d, f, nil), redactor(d, f, nil))
|
|
||||||
if len(set.merchantIDs) != 35 || len(set.tags) != 36 {
|
|
||||||
t.Fatalf("complete registry omitted entries: merchants=%d tags=%d", len(set.merchantIDs), len(set.tags))
|
|
||||||
}
|
|
||||||
if set.merchantIDs["mer_34"] != "mer_34" ||
|
|
||||||
set.tagIDs["tag_34"] != "tag_34" ||
|
|
||||||
set.categoryIDs["cat_34"] != "cat_34" {
|
|
||||||
t.Fatal("registry omitted real ids")
|
|
||||||
}
|
|
||||||
before := domain.Clone(d)
|
before := domain.Clone(d)
|
||||||
|
for _, mode := range []string{"existing", "duplicate name"} {
|
||||||
|
t.Run(mode, func(t *testing.T) {
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
content, _ := json.Marshal(map[string]any{
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
"merchant_id": "mer_34",
|
if len(prompt.Merchants) != 35 || len(prompt.Tags) != 36 || len(prompt.Categories) != 37 {
|
||||||
"new_merchant": nil,
|
t.Fatalf("complete candidates missing: merchants=%d tags=%d categories=%d", len(prompt.Merchants), len(prompt.Tags), len(prompt.Categories))
|
||||||
"category_id": "cat_34",
|
}
|
||||||
"tag_ids": []string{"tag_34"},
|
merchants, categories, tags := map[string]string{}, map[string]string{}, map[string]string{}
|
||||||
|
for _, merchant := range prompt.Merchants {
|
||||||
|
merchants[merchant.Name] = merchant.ID
|
||||||
|
}
|
||||||
|
for _, category := range prompt.Categories {
|
||||||
|
if category.Kind != "expense" {
|
||||||
|
t.Errorf("ineligible category candidate: %+v", category)
|
||||||
|
}
|
||||||
|
categories[category.Path] = category.ID
|
||||||
|
}
|
||||||
|
for _, tag := range prompt.Tags {
|
||||||
|
tags[tag.Name] = tag.ID
|
||||||
|
}
|
||||||
|
for _, merchant := range d.Merchants {
|
||||||
|
if merchants[normalize(merchant.Name)] == "" {
|
||||||
|
t.Errorf("merchant omitted: %s", merchant.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, category := range d.Categories {
|
||||||
|
if category.Kind == "expense" && category.ID != "cat_expenses" && categories[normalize(domain.CategoryPath(d, category.ID))] == "" {
|
||||||
|
t.Errorf("eligible category omitted: %s", category.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
for _, tag := range d.Tags {
|
||||||
|
if tags[normalize(tag.Name)] == "" {
|
||||||
|
t.Errorf("tag omitted: %s", tag.Name)
|
||||||
|
}
|
||||||
|
}
|
||||||
|
var merchantID, newMerchant any = merchants["z distant bakery"], nil
|
||||||
|
if mode == "duplicate name" {
|
||||||
|
merchantID, newMerchant = nil, "Z Distant Bakery"
|
||||||
|
}
|
||||||
|
content, err := json.Marshal(map[string]any{
|
||||||
|
"merchant_id": merchantID,
|
||||||
|
"new_merchant": newMerchant,
|
||||||
|
"category_id": categories[normalize(domain.CategoryPath(d, "cat_34"))],
|
||||||
|
"tag_ids": []string{tags["tag 34"]},
|
||||||
"confidence": "high",
|
"confidence": "high",
|
||||||
})
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
reply(w, string(content))
|
reply(w, string(content))
|
||||||
})
|
})
|
||||||
p, err := c.Classify(context.Background(), f, d, true)
|
p, err := c.Classify(context.Background(), f, d, true)
|
||||||
if err != nil || p.Enrichment.MerchantID != "mer_34" || p.Enrichment.CategoryID != "cat_34" || !reflect.DeepEqual(p.Enrichment.TagIDs, []string{"tag_34"}) {
|
if err != nil || p.NewMerchant != nil || p.Enrichment.MerchantID != "mer_34" || p.Enrichment.CategoryID != "cat_34" || !reflect.DeepEqual(p.Enrichment.TagIDs, []string{"tag_34"}) {
|
||||||
t.Fatalf("complete registry selection failed: %+v %v", p, err)
|
t.Fatalf("complete registry selection failed: %+v %v", p, err)
|
||||||
}
|
}
|
||||||
if !reflect.DeepEqual(before, d) {
|
if !reflect.DeepEqual(before, d) {
|
||||||
t.Fatal("retrieval mutated registry order")
|
t.Fatal("classification mutated the dataset")
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -480,7 +561,7 @@ func TestConfiguredPrivateNamesAndIdentifiersRedactWithoutRemovingPayee(t *testi
|
|||||||
func TestLowConfidenceKeepsProposalAndRecordsConfidence(t *testing.T) {
|
func TestLowConfidenceKeepsProposalAndRecordsConfidence(t *testing.T) {
|
||||||
f, d := fixture()
|
f, d := fixture()
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
reply(w, `{"merchant_id":"mer_coffee","new_merchant":null,"category_id":"cat_food","tag_ids":["tag_daily"],"confidence":"low"}`)
|
reply(w, `{"merchant_id":"m1","new_merchant":null,"category_id":"c1","tag_ids":["t1"],"confidence":"low"}`)
|
||||||
})
|
})
|
||||||
p, err := c.Classify(context.Background(), f, d, true)
|
p, err := c.Classify(context.Background(), f, d, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -547,7 +628,7 @@ func TestPayeeAndPublicMerchantAreSentToAI(t *testing.T) {
|
|||||||
Description string `json:"description"`
|
Description string `json:"description"`
|
||||||
Counterparty string `json:"counterparty"`
|
Counterparty string `json:"counterparty"`
|
||||||
} `json:"transaction"`
|
} `json:"transaction"`
|
||||||
Merchants []candidate `json:"merchants"`
|
Merchants []merchantPrompt `json:"merchants"`
|
||||||
}
|
}
|
||||||
if err := json.Unmarshal([]byte(req.Messages[1].Content), &prompt); err != nil {
|
if err := json.Unmarshal([]byte(req.Messages[1].Content), &prompt); err != nil {
|
||||||
t.Fatal(err)
|
t.Fatal(err)
|
||||||
@@ -558,7 +639,7 @@ func TestPayeeAndPublicMerchantAreSentToAI(t *testing.T) {
|
|||||||
if len(prompt.Merchants) != 26 || prompt.Merchants[0].Name != "coffee house" {
|
if len(prompt.Merchants) != 26 || prompt.Merchants[0].Name != "coffee house" {
|
||||||
t.Fatalf("complete merchant registry missing: %d", len(prompt.Merchants))
|
t.Fatalf("complete merchant registry missing: %d", len(prompt.Merchants))
|
||||||
}
|
}
|
||||||
reply(w, `{"merchant_id":"mer_coffee","new_merchant":null,"category_id":"cat_food","tag_ids":[],"confidence":"high"}`)
|
reply(w, fmt.Sprintf(`{"merchant_id":%q,"new_merchant":null,"category_id":"c1","tag_ids":[],"confidence":"high"}`, prompt.Merchants[0].ID))
|
||||||
})
|
})
|
||||||
p, err := c.Classify(context.Background(), f, d, true)
|
p, err := c.Classify(context.Background(), f, d, true)
|
||||||
if err != nil || p.Enrichment.MerchantID != "mer_coffee" {
|
if err != nil || p.Enrichment.MerchantID != "mer_coffee" {
|
||||||
|
|||||||
@@ -2,9 +2,11 @@ package classification
|
|||||||
|
|
||||||
import (
|
import (
|
||||||
"context"
|
"context"
|
||||||
|
"encoding/json"
|
||||||
"fmt"
|
"fmt"
|
||||||
"net/http"
|
"net/http"
|
||||||
"reflect"
|
"reflect"
|
||||||
|
"regexp"
|
||||||
"strings"
|
"strings"
|
||||||
"testing"
|
"testing"
|
||||||
|
|
||||||
@@ -60,6 +62,17 @@ func ledgerFixture() (domain.Dataset, domain.Facts) {
|
|||||||
return d, facts
|
return d, facts
|
||||||
}
|
}
|
||||||
|
|
||||||
|
func categoryRefForPath(t *testing.T, categories []categoryPrompt, path string) string {
|
||||||
|
t.Helper()
|
||||||
|
for _, category := range categories {
|
||||||
|
if category.Path == path {
|
||||||
|
return category.ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
t.Errorf("category path %q missing from prompt: %+v", path, categories)
|
||||||
|
return ""
|
||||||
|
}
|
||||||
|
|
||||||
// strictKeywords is what every targeted provider accepts in strict
|
// strictKeywords is what every targeted provider accepts in strict
|
||||||
// structured-output mode. uniqueItems is rejected outright by OpenAI-family
|
// structured-output mode. uniqueItems is rejected outright by OpenAI-family
|
||||||
// endpoints ("'uniqueItems' is not permitted"); minItems/maxItems make Gemini
|
// endpoints ("'uniqueItems' is not permitted"); minItems/maxItems make Gemini
|
||||||
@@ -115,7 +128,9 @@ func TestLedgerRowClassifiesThroughStrictSchema(t *testing.T) {
|
|||||||
}
|
}
|
||||||
}
|
}
|
||||||
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
reply(w, `{"merchant_id":null,"new_merchant":"Finanzamt Bruehl","category_id":"`+taxes+`","tag_ids":[],"confidence":"high"}`)
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
|
category := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, taxes)))
|
||||||
|
reply(w, `{"merchant_id":null,"new_merchant":"Finanzamt Bruehl","category_id":"`+category+`","tag_ids":[],"confidence":"high"}`)
|
||||||
})
|
})
|
||||||
p, err := c.Classify(context.Background(), facts, d, true)
|
p, err := c.Classify(context.Background(), facts, d, true)
|
||||||
if err != nil {
|
if err != nil {
|
||||||
@@ -188,18 +203,123 @@ func TestManualCorrectionsOutrankAIPrecedent(t *testing.T) {
|
|||||||
add("tx_corrected", "2026-08-01", events, "manual")
|
add("tx_corrected", "2026-08-01", events, "manual")
|
||||||
target := domain.Facts{ID: "tx_new", AccountID: "acct_kontist", BookingDate: "2026-08-30",
|
target := domain.Facts{ID: "tx_new", AccountID: "acct_kontist", BookingDate: "2026-08-30",
|
||||||
Amount: "-13.00", Currency: "EUR", Counterparty: "LVR Landesmuseum Bonn"}
|
Amount: "-13.00", Currency: "EUR", Counterparty: "LVR Landesmuseum Bonn"}
|
||||||
rows := history(target, d, func(s string) string { return normalize(s) }, 20)
|
set := retrieve("", "expense", d, nil, nil)
|
||||||
if len(rows) == 0 || rows[0].Source != "user" || rows[0].CategoryID != events {
|
rows := set.history(target, d, normalize, 20)
|
||||||
|
eventsRef := categoryRefForPath(t, set.categories, domain.CategoryPath(d, events))
|
||||||
|
if len(rows) == 0 {
|
||||||
|
t.Fatal("manual correction missing from precedent")
|
||||||
|
}
|
||||||
|
if rows[0].Source != "user" || rows[0].CategoryID != eventsRef {
|
||||||
t.Fatalf("manual correction did not lead precedent: %+v", rows[0])
|
t.Fatalf("manual correction did not lead precedent: %+v", rows[0])
|
||||||
}
|
}
|
||||||
// The correction keeps its slot even in a window the AI rows could fill.
|
}
|
||||||
users := 0
|
|
||||||
for _, row := range rows {
|
func TestHistoryReferencesResolveThroughCurrentRequestCandidates(t *testing.T) {
|
||||||
if row.Source == "user" {
|
facts, d := fixture()
|
||||||
users++
|
d.Categories = append(d.Categories, domain.Category{ID: "cat_salary", Name: "Salary", ParentID: "cat_income", Kind: "income"})
|
||||||
|
d.Merchants = append(d.Merchants, domain.Merchant{ID: "mer_payroll", Name: "Payroll", DefaultCategoryID: "cat_salary"})
|
||||||
|
manual := facts
|
||||||
|
manual.ID, manual.Fingerprint, manual.BookingDate = "tx_manual", "fp_manual", "2026-08-01"
|
||||||
|
d.Transactions = append(d.Transactions, domain.Transaction{Facts: manual, Enrichment: domain.Enrichment{
|
||||||
|
Kind: "expense", CategoryID: "cat_food", MerchantID: "mer_coffee", TagIDs: []string{"tag_daily"},
|
||||||
|
Classification: domain.Provenance{Source: "manual"},
|
||||||
|
}})
|
||||||
|
income := manual
|
||||||
|
income.ID, income.Fingerprint, income.BookingDate, income.Amount = "tx_income", "fp_income", "2026-08-31", "100.00"
|
||||||
|
d.Transactions = append(d.Transactions, domain.Transaction{Facts: income, Enrichment: domain.Enrichment{
|
||||||
|
Kind: "income", CategoryID: "cat_salary", MerchantID: "mer_payroll", TagIDs: []string{"tag_daily"},
|
||||||
|
Classification: domain.Provenance{Source: "manual"},
|
||||||
|
}})
|
||||||
|
categoryPattern := regexp.MustCompile(`^c[1-9][0-9]*$`)
|
||||||
|
merchantPattern := regexp.MustCompile(`^m[1-9][0-9]*$`)
|
||||||
|
tagPattern := regexp.MustCompile(`^t[1-9][0-9]*$`)
|
||||||
|
expectedCategories := 2
|
||||||
|
c := mockClient(t, func(w http.ResponseWriter, r *http.Request) {
|
||||||
|
prompt := decodeClassificationPrompt(t, r)
|
||||||
|
if len(prompt.Categories) != expectedCategories || len(prompt.Merchants) != len(d.Merchants) || len(prompt.Tags) != len(d.Tags) {
|
||||||
|
t.Errorf("request lost eligible registry candidates: categories=%d merchants=%d tags=%d",
|
||||||
|
len(prompt.Categories), len(prompt.Merchants), len(prompt.Tags))
|
||||||
|
}
|
||||||
|
categories := make(map[string]bool)
|
||||||
|
for _, candidate := range prompt.Categories {
|
||||||
|
if !categoryPattern.MatchString(candidate.ID) || candidate.Kind != "expense" || categories[candidate.ID] {
|
||||||
|
t.Errorf("invalid expense category reference: %+v", candidate)
|
||||||
|
}
|
||||||
|
categories[candidate.ID] = true
|
||||||
|
}
|
||||||
|
food := categoryRefForPath(t, prompt.Categories, normalize(domain.CategoryPath(d, "cat_food")))
|
||||||
|
merchants := make(map[string]bool)
|
||||||
|
coffee := ""
|
||||||
|
for _, candidate := range prompt.Merchants {
|
||||||
|
if !merchantPattern.MatchString(candidate.ID) || merchants[candidate.ID] {
|
||||||
|
t.Errorf("invalid merchant reference: %+v", candidate)
|
||||||
|
}
|
||||||
|
merchants[candidate.ID] = true
|
||||||
|
if candidate.UsualCategory != "" && !categories[candidate.UsualCategory] {
|
||||||
|
t.Errorf("merchant has dangling usual category: %+v", candidate)
|
||||||
|
}
|
||||||
|
if candidate.Name == "coffee house" {
|
||||||
|
coffee = candidate.ID
|
||||||
|
if candidate.UsualCategory != food {
|
||||||
|
t.Errorf("merchant usual category does not identify Food: %+v", candidate)
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
if users == 0 {
|
}
|
||||||
t.Fatal("correction crowded out of the history window")
|
tags := make(map[string]bool)
|
||||||
|
daily := ""
|
||||||
|
for _, candidate := range prompt.Tags {
|
||||||
|
if !tagPattern.MatchString(candidate.ID) || tags[candidate.ID] {
|
||||||
|
t.Errorf("invalid tag reference: %+v", candidate)
|
||||||
|
}
|
||||||
|
tags[candidate.ID] = true
|
||||||
|
if candidate.Name == "daily" {
|
||||||
|
daily = candidate.ID
|
||||||
|
}
|
||||||
|
}
|
||||||
|
if coffee == "" || daily == "" {
|
||||||
|
t.Error("request lost Coffee House or Daily")
|
||||||
|
}
|
||||||
|
if len(prompt.History) != 1 {
|
||||||
|
t.Errorf("expected only applicable manual expense history, got %+v", prompt.History)
|
||||||
|
w.WriteHeader(http.StatusBadRequest)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
history := prompt.History[0]
|
||||||
|
if history.Source != "user" || history.CategoryID != food || history.MerchantID != coffee ||
|
||||||
|
!reflect.DeepEqual(history.TagIDs, []string{daily}) {
|
||||||
|
t.Errorf("manual history references do not match offered records: %+v", history)
|
||||||
|
}
|
||||||
|
// Copying the correction must select the original registry records, not
|
||||||
|
// whatever records occupied these request-local references previously.
|
||||||
|
answer, err := json.Marshal(map[string]any{
|
||||||
|
"merchant_id": history.MerchantID, "new_merchant": nil,
|
||||||
|
"category_id": history.CategoryID, "tag_ids": history.TagIDs, "confidence": "high",
|
||||||
|
})
|
||||||
|
if err != nil {
|
||||||
|
t.Error(err)
|
||||||
|
w.WriteHeader(http.StatusInternalServerError)
|
||||||
|
return
|
||||||
|
}
|
||||||
|
reply(w, string(answer))
|
||||||
|
})
|
||||||
|
for _, name := range []string{"original registry", "shifted registry"} {
|
||||||
|
if name == "shifted registry" {
|
||||||
|
// New names sort before every selected record and change all three
|
||||||
|
// references without changing the canonical correction.
|
||||||
|
d.Categories = append(d.Categories, domain.Category{ID: "cat_early", Name: "Aardvark", ParentID: "cat_expenses", Kind: "expense"})
|
||||||
|
d.Merchants = append(d.Merchants, domain.Merchant{ID: "mer_early", Name: "Aardvark"})
|
||||||
|
d.Tags = append(d.Tags, domain.Tag{ID: "tag_early", Name: "Aardvark"})
|
||||||
|
expectedCategories++
|
||||||
|
}
|
||||||
|
t.Run(name, func(t *testing.T) {
|
||||||
|
p, err := c.Classify(context.Background(), facts, d, true)
|
||||||
|
if err != nil {
|
||||||
|
t.Fatal(err)
|
||||||
|
}
|
||||||
|
if p.NewMerchant != nil || p.Enrichment.CategoryID != "cat_food" || p.Enrichment.MerchantID != "mer_coffee" ||
|
||||||
|
!reflect.DeepEqual(p.Enrichment.TagIDs, []string{"tag_daily"}) {
|
||||||
|
t.Fatalf("manual precedent resolved to wrong canonical records: %+v", p)
|
||||||
|
}
|
||||||
|
})
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|||||||
Reference in New Issue
Block a user