schemas.py 2.0 KB

1234567891011121314151617181920212223242526272829303132333435363738394041424344454647484950515253545556575859606162636465666768697071727374757677787980818283
  1. from __future__ import annotations
  2. from enum import StrEnum
  3. from typing import Any
  4. from pydantic import BaseModel, Field
  5. class RouteName(StrEnum):
  6. DIRECT_ANSWER = "direct_answer"
  7. MILVUS_SEARCH = "milvus_search"
  8. SQL_QUERY = "sql_query"
  9. WEB_SEARCH = "web_search"
  10. MULTI_SOURCE = "multi_source"
  11. CLARIFY = "clarify"
  12. REFUSE = "refuse"
  13. class RouteDecision(BaseModel):
  14. needs_retrieval: bool
  15. intent: str
  16. routes: list[RouteName] = Field(default_factory=list)
  17. requires_decomposition: bool = False
  18. confidence: str = "medium"
  19. reason_code: str
  20. filters: dict[str, Any] = Field(default_factory=dict)
  21. max_rounds: int = 2
  22. fallback: str = "clarify"
  23. class PlanStep(BaseModel):
  24. id: str
  25. tool: RouteName
  26. query: str
  27. arguments: dict[str, Any] = Field(default_factory=dict)
  28. depends_on: list[str] = Field(default_factory=list)
  29. class RetrievalPlan(BaseModel):
  30. goal: str
  31. steps: list[PlanStep]
  32. class Evidence(BaseModel):
  33. source_type: str
  34. source: str
  35. content: str
  36. score: float | None = None
  37. metadata: dict[str, Any] = Field(default_factory=dict)
  38. class ToolResult(BaseModel):
  39. status: str
  40. tool: str
  41. data: list[dict[str, Any]] = Field(default_factory=list)
  42. evidence: list[Evidence] = Field(default_factory=list)
  43. latency_ms: int = 0
  44. retryable: bool = False
  45. error_code: str | None = None
  46. error_message: str | None = None
  47. class QualityGrade(BaseModel):
  48. relevant: bool
  49. sufficient: bool
  50. missing_aspects: list[str] = Field(default_factory=list)
  51. recommended_action: str
  52. reason: str
  53. class QueryRequest(BaseModel):
  54. query: str = Field(min_length=1, max_length=4000)
  55. session_id: str = "default"
  56. debug: bool = False
  57. class QueryResponse(BaseModel):
  58. answer: str
  59. citations: list[Evidence]
  60. route: RouteDecision
  61. executed_queries: list[dict[str, Any]] = Field(default_factory=list)
  62. trace: list[dict[str, Any]] = Field(default_factory=list)
  63. termination_reason: str