diff --git a/backend/ai.py b/backend/ai.py index d16befa..5bfff01 100644 --- a/backend/ai.py +++ b/backend/ai.py @@ -8,7 +8,7 @@ from anthropic import Anthropic from dotenv import load_dotenv -from database import db_get +from database import db_get, escape_filter_value load_dotenv() @@ -28,14 +28,16 @@ def _get_client() -> Anthropic: def _fetch_ticket(ticket_id: str) -> dict[str, Any]: - result = db_get("tickets", f"id=eq.{ticket_id}") + result = db_get("tickets", f"id=eq.{escape_filter_value(ticket_id)}") if not isinstance(result, list) or not result: raise ValueError(f"Ticket {ticket_id} not found") return result[0] def _fetch_notes(ticket_id: str) -> list[dict[str, Any]]: - notes = db_get("ticket_notes", f"ticket_id=eq.{ticket_id}&order=created_at.asc") + notes = db_get( + "ticket_notes", f"ticket_id=eq.{escape_filter_value(ticket_id)}&order=created_at.asc" + ) return notes if isinstance(notes, list) else [] diff --git a/backend/test_ai.py b/backend/test_ai.py new file mode 100644 index 0000000..b5ea720 --- /dev/null +++ b/backend/test_ai.py @@ -0,0 +1,46 @@ +from unittest.mock import MagicMock + +import pytest + +import ai +import database + + +@pytest.fixture +def fake_supabase(monkeypatch): + monkeypatch.setattr(database, "SUPABASE_URL", "https://example.supabase.co") + monkeypatch.setattr(database, "SUPABASE_KEY", "secret") + fake_response = MagicMock() + fake_response.status_code = 200 + fake_response.text = "[]" + fake_response.json.return_value = [] + fake_request = MagicMock(return_value=fake_response) + monkeypatch.setattr(database.requests, "request", fake_request) + return fake_request + + +# --------------------------------------------------------------------------- +# _fetch_ticket / _fetch_notes -- ticket_id is escaped before interpolation +# into PostgREST filters, mirroring the guard already applied in main.py. +# --------------------------------------------------------------------------- + + +def test_fetch_ticket_escapes_postgrest_injection(fake_supabase): + malicious_id = "t1&select=*,customers(email)" + + with pytest.raises(ValueError, match="not found"): + ai._fetch_ticket(malicious_id) + + url = fake_supabase.call_args.args[1] + assert url.count("&") == 0, url + + +def test_fetch_notes_escapes_postgrest_injection(fake_supabase): + malicious_id = "t1&status=eq.Closed" + + notes = ai._fetch_notes(malicious_id) + + assert notes == [] + url = fake_supabase.call_args.args[1] + assert url.count("&") == 1, url + assert "order=created_at.asc" in url