Spaces:
Sleeping
Sleeping
| """Pydantic contracts for the local SAGE service.""" | |
| from __future__ import annotations | |
| from typing import Any, Literal | |
| from pydantic import BaseModel, Field, model_validator | |
| StageNumber = Literal[1, 2, 3] | |
| RunState = Literal["queued", "running", "succeeded", "failed", "cancel_requested", "cancelled"] | |
| StageState = Literal["pending", "running", "succeeded", "failed", "skipped", "cancelled"] | |
| class RunRequest(BaseModel): | |
| """Request body for creating a pipeline run.""" | |
| input_path: str | None = Field(default=None, description="Repo-relative or absolute markdown input path.") | |
| request_text: str | None = Field(default=None, description="Inline request text used when no input_path is supplied.") | |
| case_id: str | None = Field(default=None, description="Human-readable case identifier.") | |
| run_id: str | None = Field(default=None, description="Optional stable run id. Must be path-safe.") | |
| start_stage: StageNumber = 1 | |
| end_stage: StageNumber = 3 | |
| max_llm_calls: int = Field(default=1000, ge=1, le=5000) | |
| llm_dotenv_path: str | None = ".env" | |
| llm_provider: str | None = Field( | |
| default=None, | |
| description="Optional provider override, such as deepseek, azure, local-vllm, or qwen.", | |
| ) | |
| llm_model: str | None = Field(default=None, description="Optional provider-specific model/deployment id.") | |
| llm_response_format: str | None = Field( | |
| default=None, | |
| description="Optional response_format mode override: json_schema, json_object, or none.", | |
| ) | |
| temperature: float = Field(default=0.0, ge=0.0, le=2.0) | |
| llm_max_retries: int = Field(default=2, ge=0, le=10) | |
| timeout_seconds: float = Field(default=90.0, ge=1.0, le=600.0) | |
| require_human_confirmation: bool = Field( | |
| default=False, | |
| description="When true, the service pauses after Stage 2 so users can verify non-metadata candidates before SQL generation.", | |
| ) | |
| stage2_top_k: int = Field(default=10, ge=1, le=100) | |
| stage2_mode: str = Field( | |
| default="adopt_suggestions", | |
| description="adopt_suggestions: auto-accept LLM-cleaned evidence for Stage 3. human_review: write suggestions but wait for human confirmation.", | |
| ) | |
| stage2_llm_verify_results: bool = False | |
| stage2_llm_filter_noise: bool = False | |
| stage2_llm_retry_limit: int = Field(default=1, ge=0, le=5) | |
| ir_json_path: str | None = Field(default=None, description="Existing Stage 1 IR path for runs starting at Stage 2/3.") | |
| retrieval_context_path: str | None = Field( | |
| default=None, | |
| description="Existing Stage 2 context path for runs starting at Stage 3.", | |
| ) | |
| def validate_stage_inputs(self) -> "RunRequest": | |
| if self.end_stage < self.start_stage: | |
| raise ValueError("end_stage must be greater than or equal to start_stage.") | |
| if self.start_stage == 1 and not self.input_path and not self.request_text: | |
| raise ValueError("Stage 1 runs require input_path or request_text.") | |
| if self.start_stage >= 2 and not self.ir_json_path: | |
| # The server can also reuse artifacts/<run_id>/stage_01/cohortbuild_ir.json | |
| # when the caller supplies a run_id. The pipeline validates that at runtime. | |
| pass | |
| return self | |
| class StageRecord(BaseModel): | |
| name: str | |
| state: StageState = "pending" | |
| started_at: str | None = None | |
| finished_at: str | None = None | |
| output_dir: str | None = None | |
| llm_calls: int = 0 | |
| summary: dict[str, Any] = Field(default_factory=dict) | |
| error: str | None = None | |
| class ArtifactInfo(BaseModel): | |
| name: str | |
| path: str | |
| size_bytes: int | |
| url: str | |
| class RunStatus(BaseModel): | |
| run_id: str | |
| case_id: str | |
| state: RunState | |
| created_at: str | |
| updated_at: str | |
| input_path: str | None = None | |
| run_dir: str | |
| start_stage: StageNumber | |
| end_stage: StageNumber | |
| llm_calls: int = 0 | |
| stages: dict[str, StageRecord] = Field(default_factory=dict) | |
| artifacts: dict[str, str] = Field(default_factory=dict) | |
| artifact_count: int = 0 | |
| events: list[dict[str, Any]] = Field(default_factory=list) | |
| error: str | None = None | |