from fastapi import FastAPI from fastapi.testclient import TestClient from api.dependencies import get_workflow_manager from api.endpoints import router from workflows.workflow_manager import WorkflowType class GuardWorkflowManager: def execute_workflow(self, *args, **kwargs): raise AssertionError("execute_workflow should not be called for invalid query payloads") class StubWorkflowManager: def __init__(self): self.calls = [] def execute_workflow(self, workflow_type, user_input, session_id=None, **kwargs): self.calls.append( { "workflow_type": workflow_type, "user_input": user_input, "session_id": session_id, "kwargs": kwargs, } ) return { "session_id": session_id or "cid-1", "workflow_type": workflow_type.value if isinstance(workflow_type, WorkflowType) else str(workflow_type), "result": {"context": {"final_sql": "SELECT 1"}, "ok": True}, } class DisabledMessageStorage: enabled = False def save_message(self, **kwargs): return False class CaptureMessageStorage: enabled = True def __init__(self): self.saved = [] self.created = [] self.existing = {} def create_conversation(self, conversation_id, user, name, status, introduction, created_at, updated_at): record = { "conversation_id": conversation_id, "user": user, "name": name, "status": status, "introduction": introduction, "created_at": created_at, "updated_at": updated_at, } self.created.append(record) self.existing[conversation_id] = record return True def get_conversation_by_id(self, conversation_id): return self.existing.get(conversation_id) def update_conversation_updated_at(self, conversation_id, updated_at): if conversation_id not in self.existing: return False self.existing[conversation_id]["updated_at"] = updated_at return True def save_message(self, **kwargs): self.saved.append(kwargs) return True def _build_client(workflow_manager) -> TestClient: app = FastAPI() app.include_router(router) app.dependency_overrides[get_workflow_manager] = lambda: workflow_manager return TestClient(app) def _payload(query, conversation_id="cid-123"): return { "query": query, "inputs": {}, "response_mode": "streaming", "user": "tester", "conversation_id": conversation_id, "files": [], } def test_invalid_query_returns_400_for_stream_endpoint(monkeypatch): monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage()) client = _build_client(GuardWorkflowManager()) for invalid_query in (None, "", " "): response = client.post("/api/workflows/stream", json=_payload(invalid_query)) assert response.status_code == 400 assert response.json() == { "detail": { "code": "INVALID_REQUEST", "message": "query 不能为空", "detail": {"field": "query", "reason": "missing_or_blank"}, } } def test_invalid_query_returns_400_for_blocking_endpoints(monkeypatch): monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage()) client = _build_client(GuardWorkflowManager()) for path in ("/api/workflows", "/api/sql/generate"): response = client.post(path, json=_payload(" ")) assert response.status_code == 400 assert response.json()["detail"]["code"] == "INVALID_REQUEST" assert response.json()["detail"]["detail"] == {"field": "query", "reason": "missing_or_blank"} def test_valid_query_still_reaches_workflow_manager(monkeypatch): monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage()) workflow_manager = StubWorkflowManager() client = _build_client(workflow_manager) response = client.post( "/api/workflows", json={ **_payload("查询 SO 4020438779 的 eta 信息"), "response_mode": "blocking", }, ) assert response.status_code == 200 assert workflow_manager.calls[0]["user_input"] == "查询 SO 4020438779 的 eta 信息" assert response.json()["session_id"] == "cid-123" assert response.json()["workflow_type"] == "conversation" def test_valid_query_persists_message_record(monkeypatch): storage = CaptureMessageStorage() monkeypatch.setattr("api.endpoints.get_message_storage", lambda: storage) workflow_manager = StubWorkflowManager() client = _build_client(workflow_manager) response = client.post( "/api/workflows", json={ **_payload("查询 SO 4020438779 的 eta 信息", conversation_id=None), "response_mode": "blocking", }, ) assert response.status_code == 200 assert len(storage.saved) == 1 saved = storage.saved[0] assert saved["conversation_id"] == response.json()["session_id"] assert saved["query"] == "查询 SO 4020438779 的 eta 信息" assert saved["workflow_type"] == "conversation" assert saved["created_at"] == saved["updated_at"] assert saved["logs"][0].startswith("run_workflow.start") assert any(item.startswith("run_workflow.success") for item in saved["logs"]) assert storage.created[0]["name"] == "查询 SO 4020438779 的 e" def test_missing_conversation_id_returns_400_for_blocking_workflow(monkeypatch): monkeypatch.setattr("api.endpoints.get_message_storage", lambda: CaptureMessageStorage()) workflow_manager = StubWorkflowManager() client = _build_client(workflow_manager) response = client.post( "/api/workflows", json={ **_payload("查询 SO 4020438779 的 eta 信息", conversation_id="cid-missing"), "response_mode": "blocking", }, ) assert response.status_code == 400 assert response.json()["detail"]["code"] == "CONVERSATION_NOT_FOUND" def test_invalid_response_mode_uses_dedicated_error_code(monkeypatch): monkeypatch.setattr("api.endpoints.get_message_storage", lambda: DisabledMessageStorage()) client = _build_client(GuardWorkflowManager()) response = client.post( "/api/workflows/stream", json={ **_payload("查询 SO 4020438779 的 eta 信息"), "response_mode": "blocking", }, ) assert response.status_code == 400 assert response.json()["detail"]["code"] == "INVALID_RESPONSE_MODE"