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 := `- 1
- 2
`
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)
}
}