"""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.", ) @model_validator(mode="after") 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//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