AI_Agro_Support/pipeline/config.py
Arsham Mirehvandi 099d1e0af9 Initial commit
2026-08-22 08:47:56 +02:00

235 lines
7.2 KiB
Python

"""Load and merge config.yaml with environment variables."""
from __future__ import annotations
import os
from dataclasses import dataclass
from pathlib import Path
import yaml
from dotenv import load_dotenv
ROOT = Path(__file__).resolve().parent.parent
@dataclass(frozen=True)
class LlmSettings:
provider: str
openai_model: str
anthropic_model: str
gemini_model: str
max_tokens: int
temperature: float
@dataclass(frozen=True)
class SqlSettings:
driver: str
server: str
database: str
username: str
password: str
@dataclass(frozen=True)
class WeaviateSettings:
host: str
http_port: int
grpc_port: int
@dataclass(frozen=True)
class DiseasePlan:
"""One crop-disease pair to run the batch for."""
model_name: str # matched against AI_agrosupport_agro_models.anmod_model
disease: str # canonical English disease name (vocab/diseases.yaml)
@dataclass(frozen=True)
class CropPlan:
"""A crop and the diseases to check for every field growing it."""
crop: str # canonical English crop name (vocab/crops.yaml)
diseases: tuple[DiseasePlan, ...]
@dataclass(frozen=True)
class WorklistSettings:
require_enabled_model: bool
field_allowlist: tuple[int, ...]
@dataclass(frozen=True)
class ConcurrencySettings:
workers: int
sql_limit: int
gemini_limit: int
gemini_requests_per_minute: int
@dataclass(frozen=True)
class RetrySettings:
attempts: int
initial_backoff_seconds: float
max_backoff_seconds: float
@dataclass(frozen=True)
class ScheduleSettings:
deadline: str # "HH:MM", local time
@dataclass(frozen=True)
class Settings:
root: Path
field_id: int | None
disease_name: str | None
crops: tuple[CropPlan, ...]
worklist: WorklistSettings
concurrency: ConcurrencySettings
retry: RetrySettings
schedule: ScheduleSettings
crops_vocab: Path
diseases_vocab: Path
llm_query_synthesis: LlmSettings
llm_advice_generation: LlmSettings
sql: SqlSettings
weaviate: WeaviateSettings
gemini_api_key: str
openai_api_key: str
anthropic_api_key: str
prompts_dir: Path
output_dir: Path
def _parse_llm_settings(raw: dict, defaults: dict) -> LlmSettings:
"""Build an LlmSettings profile from a config sub-block, falling back to defaults."""
return LlmSettings(
provider=str(raw.get("provider", defaults["provider"])).strip().lower(),
openai_model=str(raw.get("openai_model", defaults["openai_model"])),
anthropic_model=str(raw.get("anthropic_model", defaults["anthropic_model"])),
gemini_model=str(raw.get("gemini_model", defaults["gemini_model"])),
max_tokens=int(raw.get("max_tokens", defaults["max_tokens"])),
temperature=float(raw.get("temperature", defaults["temperature"])),
)
def _parse_crops(raw: list[dict] | None) -> tuple[CropPlan, ...]:
plans: list[CropPlan] = []
for entry in raw or []:
crop = str(entry["crop"]).strip()
diseases = tuple(
DiseasePlan(
model_name=str(d["model_name"]).strip(),
disease=str(d["disease"]).strip(),
)
for d in entry.get("diseases", [])
)
plans.append(CropPlan(crop=crop, diseases=diseases))
return tuple(plans)
def _parse_worklist(raw: dict) -> WorklistSettings:
return WorklistSettings(
require_enabled_model=bool(raw.get("require_enabled_model", True)),
field_allowlist=tuple(int(v) for v in raw.get("field_allowlist", []) or []),
)
def _parse_concurrency(raw: dict) -> ConcurrencySettings:
limits = raw.get("limits", {}) or {}
return ConcurrencySettings(
workers=int(raw.get("workers", 8)),
sql_limit=int(limits.get("sql", 6)),
gemini_limit=int(limits.get("gemini", 4)),
gemini_requests_per_minute=int(raw.get("gemini_requests_per_minute", 60)),
)
def _parse_retry(raw: dict) -> RetrySettings:
return RetrySettings(
attempts=int(raw.get("attempts", 3)),
initial_backoff_seconds=float(raw.get("initial_backoff_seconds", 2)),
max_backoff_seconds=float(raw.get("max_backoff_seconds", 30)),
)
def _parse_schedule(raw: dict) -> ScheduleSettings:
return ScheduleSettings(deadline=str(raw.get("deadline", "09:00")))
def load_settings(
config_path: Path | None = None,
field_id_override: int | None = None,
disease_name_override: str | None = None,
) -> Settings:
"""Load YAML config and .env into a Settings object."""
load_dotenv(ROOT / ".env")
path = config_path or (ROOT / "config.yaml")
with path.open(encoding="utf-8") as fh:
raw = yaml.safe_load(fh)
llm_raw = raw.get("llm", {})
vocab = raw.get("vocab", {})
# field_id / disease_name only matter for `python -m pipeline one`; the
# batch worklist is driven entirely by the crops: block below.
field_id = field_id_override if field_id_override is not None else raw.get("field_id")
field_id = int(field_id) if field_id is not None else None
disease_name = disease_name_override or raw.get("disease_name")
disease_name = str(disease_name).strip() if disease_name is not None else None
query_synthesis_defaults = {
"provider": "gemini",
"openai_model": "gpt-4o-mini",
"anthropic_model": "claude-haiku-4-5",
"gemini_model": "gemini-2.5-flash",
"max_tokens": 8192,
"temperature": 0.0,
}
advice_generation_defaults = {
"provider": "gemini",
"openai_model": "gpt-4o",
"anthropic_model": "claude-opus-4-5",
"gemini_model": "gemini-2.5-pro",
"max_tokens": 65536,
"temperature": 0.0,
}
return Settings(
root=ROOT,
field_id=field_id,
disease_name=disease_name,
crops=_parse_crops(raw.get("crops")),
worklist=_parse_worklist(raw.get("worklist", {})),
concurrency=_parse_concurrency(raw.get("concurrency", {})),
retry=_parse_retry(raw.get("retry", {})),
schedule=_parse_schedule(raw.get("schedule", {})),
crops_vocab=ROOT / vocab.get("crops", "vocab/crops.yaml"),
diseases_vocab=ROOT / vocab.get("diseases", "vocab/diseases.yaml"),
llm_query_synthesis=_parse_llm_settings(
llm_raw.get("query_synthesis", {}), query_synthesis_defaults
),
llm_advice_generation=_parse_llm_settings(
llm_raw.get("advice_generation", {}), advice_generation_defaults
),
sql=SqlSettings(
driver=os.environ.get("SQL_DRIVER", "ODBC Driver 17 for SQL Server"),
server=os.environ["SQL_SERVER"],
database=os.environ["SQL_DATABASE"],
username=os.environ["SQL_USERNAME"],
password=os.environ["SQL_PASSWORD"],
),
weaviate=WeaviateSettings(
host=os.environ.get("WEAVIATE_HOST", "localhost"),
http_port=int(os.environ.get("WEAVIATE_HTTP_PORT", "8080")),
grpc_port=int(os.environ.get("WEAVIATE_GRPC_PORT", "50051")),
),
gemini_api_key=os.environ.get("GEMINI_API_KEY", ""),
openai_api_key=os.environ.get("OPENAI_API_KEY", ""),
anthropic_api_key=os.environ.get("ANTHROPIC_API_KEY", ""),
prompts_dir=ROOT / "prompts",
output_dir=ROOT / "output",
)