package feeds import ( "context" "errors" "net/http" "net/http/httptest" "net/url" "os" "strings" "testing" "github.com/google/uuid" ) func TestParseCSVAndMappings(t *testing.T) { csv := "EAN,Title\n123,Widget\n456,\n" mappings := parseMappings([]any{ map[string]any{"source": "EAN", "target": "gtin"}, map[string]any{"column": "Title", "fieldName": "title"}, }) var rows []map[string]any n, err := parseCSV(strings.NewReader(csv), func(row feedRow) error { mapped, gtin := applyMappings(row, mappings) rows = append(rows, map[string]any{"gtin": gtin, "mapped": mapped}) return nil }) if err != nil { t.Fatal(err) } if n != 2 { t.Fatalf("rows=%d", n) } if rows[0]["gtin"] != "123" { t.Fatalf("gtin=%v", rows[0]["gtin"]) } m := rows[0]["mapped"].(map[string]any) if m["title"] != "Widget" { t.Fatalf("mapped=%v", m) } } func TestParseXMLItems(t *testing.T) { xmlBody := `999X` mappings := parseMappings(map[string]any{ "gtin": map[string]any{"fieldName": "gtin"}, "title": map[string]any{"fieldName": "title"}, }) var got string n, err := parseXMLItems(strings.NewReader(xmlBody), "item", func(row feedRow) error { _, gtin := applyMappings(row, mappings) got = gtin return nil }) if err != nil { t.Fatal(err) } if n != 1 || got != "999" { t.Fatalf("n=%d gtin=%q", n, got) } } func TestDownloadRejectsPrivateAndFTP(t *testing.T) { ctx := context.Background() if _, err := downloadFeed(ctx, "ftp://example.com/a.csv"); err == nil { t.Fatal("expected ftp error") } if _, err := downloadFeed(ctx, "http://127.0.0.1/x"); err == nil { t.Fatal("expected private IP error") } if _, err := downloadFeed(ctx, "http://localhost/x"); err == nil { t.Fatal("expected localhost error") } } func TestSSRFTransportDisablesEnvProxy(t *testing.T) { tr := ssrfTransport() if tr.Proxy != nil { t.Fatal("feed SSRF transport must not use ProxyFromEnvironment") } } func TestValidateFeedURLRejectsPrivateAndFTP(t *testing.T) { ctx := context.Background() if err := ValidateFeedURL(ctx, ""); err != nil { t.Fatalf("empty url should be ok: %v", err) } if err := ValidateFeedURL(ctx, "ftp://example.com/a.csv"); err == nil { t.Fatal("expected ftp error") } if err := ValidateFeedURL(ctx, "http://127.0.0.1/x"); err == nil { t.Fatal("expected private IP error") } if err := ValidateFeedURL(ctx, "http://169.254.169.254/latest"); err == nil { t.Fatal("expected metadata IP error") } if err := ValidateFeedURL(ctx, "https://8.8.8.8/feed.xml"); err != nil { t.Fatalf("public IP https should be ok: %v", err) } } func TestDownloadPublicOKWithSizeCap(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/csv") _, _ = w.Write([]byte("EAN,Title\n1,A\n")) })) t.Cleanup(srv.Close) // httptest uses 127.0.0.1 — should be blocked by SSRF guard. _, err := downloadFeed(context.Background(), srv.URL) if err == nil || !strings.Contains(err.Error(), "private") && err != errURLPrivate { // allow either wrapped or direct if err == nil { t.Fatal("expected loopback blocked") } } } func TestDownloadAllowlistPrivateHost(t *testing.T) { srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/csv") _, _ = w.Write([]byte("EAN,Title\n1,A\n")) })) t.Cleanup(srv.Close) u, err := url.Parse(srv.URL) if err != nil { t.Fatal(err) } if err := ConfigurePrivateAllowlist([]string{u.Hostname()}, []string{"127.0.0.0/8"}); err != nil { t.Fatal(err) } t.Cleanup(func() { _ = ConfigurePrivateAllowlist(nil, nil) }) blob, err := downloadFeed(context.Background(), srv.URL) if err != nil { t.Fatalf("allowlisted download: %v", err) } t.Cleanup(func() { _ = blob.Close() }) data, err := os.ReadFile(blob.path) if err != nil { t.Fatal(err) } if !strings.Contains(string(data), "EAN") { t.Fatalf("body=%q ct=%s", data, blob.contentType) } } func TestDownloadStreamsToTempFile(t *testing.T) { payload := "EAN,Title\n1,A\n2,B\n" srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/csv") _, _ = w.Write([]byte(payload)) })) t.Cleanup(srv.Close) u, err := url.Parse(srv.URL) if err != nil { t.Fatal(err) } if err := ConfigurePrivateAllowlist([]string{u.Hostname()}, []string{"127.0.0.0/8"}); err != nil { t.Fatal(err) } t.Cleanup(func() { _ = ConfigurePrivateAllowlist(nil, nil) }) blob, err := downloadFeed(context.Background(), srv.URL) if err != nil { t.Fatal(err) } if !blob.owned || blob.path == "" { t.Fatalf("expected owned temp path, got %+v", blob) } if _, err := os.Stat(blob.path); err != nil { t.Fatalf("temp missing: %v", err) } hash, err := sha256HexFile(blob) if err != nil { t.Fatal(err) } wantHash := sha256Hex([]byte(payload)) if hash != wantHash { t.Fatalf("hash=%s want=%s", hash, wantHash) } f, err := blob.Open() if err != nil { t.Fatal(err) } var rows int n, err := parseCSV(f, func(feedRow) error { rows++ return nil }) _ = f.Close() if err != nil { t.Fatal(err) } if n != 2 || rows != 2 { t.Fatalf("n=%d rows=%d", n, rows) } path := blob.path if err := blob.Close(); err != nil { t.Fatal(err) } if _, err := os.Stat(path); !os.IsNotExist(err) { t.Fatalf("temp should be removed after Close, err=%v", err) } } func TestDetectFeedFormat(t *testing.T) { if detectFeedFormat("csv", "", "", nil) != "csv" { t.Fatal("csv") } if detectFeedFormat("", "application/xml", "", []byte("")) != "xml" { t.Fatal("xml") } } func TestClassifyUpsertOpsInsertTouchUpdate(t *testing.T) { idTouch := uuid.MustParse("11111111-1111-1111-1111-111111111111") idUpdate := uuid.MustParse("22222222-2222-2222-2222-222222222222") existing := map[string]existingProduct{ "touch": {ID: idTouch, MappedData: []byte(`{"title":"Same"}`)}, "update": {ID: idUpdate, MappedData: []byte(`{"title":"Old"}`)}, } chunk := []pendingProduct{ {GTIN: "new", RawData: map[string]string{"EAN": "new"}, MappedData: map[string]any{"title": "N"}}, {GTIN: "touch", RawData: map[string]string{"EAN": "touch"}, MappedData: map[string]any{"title": "Same"}}, {GTIN: "update", RawData: map[string]string{"EAN": "update"}, MappedData: map[string]any{"title": "New"}}, } ops, skipped := classifyUpsertOps(chunk, existing) if skipped != 0 { t.Fatalf("skipped=%d", skipped) } if len(ops) != 3 { t.Fatalf("ops=%d", len(ops)) } if ops[0].kind != upsertOpInsert || ops[0].gtin != "new" { t.Fatalf("op0=%+v", ops[0]) } if ops[1].kind != upsertOpTouch || ops[1].id != idTouch { t.Fatalf("op1=%+v", ops[1]) } if ops[2].kind != upsertOpUpdate || ops[2].id != idUpdate { t.Fatalf("op2=%+v", ops[2]) } inserts, touches, updates := partitionUpsertOps(ops) if len(inserts) != 1 || len(touches) != 1 || len(updates) != 1 { t.Fatalf("partition inserts=%d touches=%d updates=%d", len(inserts), len(touches), len(updates)) } if inserts[0].gtin != "new" || touches[0].id != idTouch || updates[0].id != idUpdate { t.Fatalf("partition payloads insert=%+v touch=%+v update=%+v", inserts[0], touches[0], updates[0]) } } func TestDedupePendingByGTINLastWins(t *testing.T) { chunk := []pendingProduct{ {GTIN: "a", MappedData: map[string]any{"title": "first"}}, {GTIN: "b", MappedData: map[string]any{"title": "only"}}, {GTIN: "a", MappedData: map[string]any{"title": "last"}}, } got := dedupePendingByGTIN(chunk) if len(got) != 2 { t.Fatalf("len=%d", len(got)) } if got[0].GTIN != "a" || got[0].MappedData["title"] != "last" { t.Fatalf("got[0]=%+v", got[0]) } if got[1].GTIN != "b" { t.Fatalf("got[1]=%+v", got[1]) } if dedupePendingByGTIN(nil) != nil { t.Fatal("nil in") } single := []pendingProduct{{GTIN: "x"}} if out := dedupePendingByGTIN(single); len(out) != 1 || out[0].GTIN != "x" { t.Fatalf("single=%+v", out) } } func TestUpsertSQLIsSetBased(t *testing.T) { for name, sql := range map[string]string{ "insert": upsertSQLInsertSet, "touch": upsertSQLTouchSet, "update": upsertSQLUpdateSet, } { lower := strings.ToLower(sql) switch name { case "insert", "update": if !strings.Contains(lower, "unnest(") { t.Fatalf("%s missing unnest: %s", name, sql) } case "touch": if !strings.Contains(lower, "any(") { t.Fatalf("touch missing ANY: %s", sql) } } if strings.Contains(lower, "values ($1") { t.Fatalf("%s still per-row VALUES form", name) } } } func TestParseCSVRespectsMaxRows(t *testing.T) { old := maxParseRows maxParseRows = 2 t.Cleanup(func() { maxParseRows = old }) body := "EAN\n1\n2\n3\n" n, err := parseCSV(strings.NewReader(body), func(feedRow) error { return nil }) if err == nil || !errors.Is(err, errParseTooManyRows) { t.Fatalf("n=%d err=%v want errParseTooManyRows", n, err) } if !strings.Contains(err.Error(), "(2)") { t.Fatalf("err=%v want limit in message", err) } if msg, ok := ClientError(err); !ok || !strings.Contains(msg, "max row limit") { t.Fatalf("ClientError=%q ok=%v", msg, ok) } if n != 3 { t.Fatalf("count=%d want 3 (exceeded after 3rd)", n) } } func TestParseXMLRespectsMaxRows(t *testing.T) { old := maxParseRows maxParseRows = 1 t.Cleanup(func() { maxParseRows = old }) body := `12` n, err := parseXMLItems(strings.NewReader(body), "item", func(feedRow) error { return nil }) if err == nil || !errors.Is(err, errParseTooManyRows) { t.Fatalf("n=%d err=%v want errParseTooManyRows", n, err) } if n != 2 { t.Fatalf("count=%d want 2", n) } } func TestDownloadRejectsOversizedBody(t *testing.T) { old := maxDownloadBytes maxDownloadBytes = 64 t.Cleanup(func() { maxDownloadBytes = old }) srv := httptest.NewServer(http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) { w.Header().Set("Content-Type", "text/csv") _, _ = w.Write([]byte(strings.Repeat("x", int(maxDownloadBytes)+2))) })) t.Cleanup(srv.Close) u, err := url.Parse(srv.URL) if err != nil { t.Fatal(err) } if err := ConfigurePrivateAllowlist([]string{u.Hostname()}, []string{"127.0.0.0/8"}); err != nil { t.Fatal(err) } t.Cleanup(func() { _ = ConfigurePrivateAllowlist(nil, nil) }) _, err = downloadFeed(context.Background(), srv.URL) if !errors.Is(err, errDownloadTooLarge) { t.Fatalf("err=%v want errDownloadTooLarge", err) } // Tiny test caps report bytes; production (≥1 MiB) reports MiB. if !strings.Contains(err.Error(), "max 64 bytes") { t.Fatalf("err=%v want max bytes in message", err) } if msg, ok := ClientError(err); !ok || !strings.Contains(msg, "size limit") { t.Fatalf("ClientError=%q ok=%v", msg, ok) } } func TestDefaultFeedCaps(t *testing.T) { const wantDownloadBytes int64 = 256 << 20 // 256 MiB — within 200–500 MiB band const wantParseRows = 1_000_000 if defaultMaxDownloadBytes != wantDownloadBytes { t.Fatalf("defaultMaxDownloadBytes=%d want %d", defaultMaxDownloadBytes, wantDownloadBytes) } if defaultMaxParseRows != wantParseRows { t.Fatalf("defaultMaxParseRows=%d want %d", defaultMaxParseRows, wantParseRows) } if defaultMaxDownloadBytes < 200<<20 || defaultMaxDownloadBytes > 500<<20 { t.Fatalf("defaultMaxDownloadBytes=%d outside 200–500 MiB guidance", defaultMaxDownloadBytes) } }