from __future__ import annotations

from pydantic import AliasChoices, Field, field_validator
from pydantic_settings import BaseSettings, SettingsConfigDict

LOCAL_ENVIRONMENTS = frozenset({"local", "development", "dev"})
SUPPORTED_EMBEDDING_PROVIDERS = frozenset({"openai", "ollama"})


class Settings(BaseSettings):
    """Runtime configuration for the embedding engine (API, Celery, Redis consumers)."""

    model_config = SettingsConfigDict(env_file=".env", extra="ignore")

    # Application
    APP_NAME: str = "Embedding Engine"
    ENVIRONMENT: str = Field(
        default="development",
        validation_alias=AliasChoices("ENVIRONMENT", "ENV"),
    )
    PORT: int = 6004
    LOG_LEVEL: str = "info"

    # MongoDB — own collections (embedding_records, embedding_runs)
    MONGO_URI: str = "mongodb://localhost:27017"
    DATABASE_NAME: str = "adpilot_indexing"
    # Upstream collections may live in separate databases (same MONGO_URI).
    DOCS_MONGODB_DATABASE: str = "repo_docs"
    CODE_CHUNKS_MONGODB_DATABASE: str = "adpilot_code_parser"
    COMMIT_ANALYSES_MONGODB_DATABASE: str = "adpilot_commit_intel"
    # Shared indexing run progress (repo-sync database)
    INDEXING_RUNS_ENABLED: bool = True
    INDEXING_RUNS_DATABASE: str = "adpilot_repo_sync"

    # OpenAI / embeddings
    OPENAI_API_KEY: str = ""
    EMBEDDING_MODEL: str = "text-embedding-3-small"
    EMBEDDING_PROVIDER: str = "openai"
    EMBEDDING_DIMENSION: int = 1536
    EMBEDDING_BATCH_SIZE: int = 100
    # OpenAI rejects any single embeddings request over 300k tokens. Cap the
    # per-request token budget safely below that so a batch of large chunks is
    # split across requests instead of failing the whole snapshot's embedding.
    EMBEDDING_MAX_TOKENS_PER_REQUEST: int = 250_000
    # Flush embedded vectors to Qdrant every N records instead of holding the
    # whole snapshot in memory until one final upsert. A six-figure-record repo
    # would otherwise accumulate ~1GB+ of vectors and exhaust the worker before
    # anything is persisted; incremental flushing also lands partial progress.
    EMBEDDING_UPSERT_FLUSH_SIZE: int = 1000
    OPENAI_REQUEST_TIMEOUT_SEC: int = 60
    OPENAI_EMBEDDING_BATCH_DELAY_MS: int = 100

    # Ollama (local embeddings, no API key required)
    OLLAMA_URL: str = "http://ollama:11434"

    # Qdrant
    QDRANT_URL: str = "http://localhost:6333"
    QDRANT_COLLECTION_NAME: str = "adpilot_embeddings"
    QDRANT_API_KEY: str = ""

    # Retrieval (POST /internal/retrieve)
    RETRIEVE_DEFAULT_TOP_K: int = 20
    RETRIEVE_DEFAULT_SCORE_THRESHOLD: float = 0.25
    RETRIEVE_MAX_TOP_K: int = 100
    RETRIEVE_GRAPH_EXPAND_MAX_NODES: int = 10
    RETRIEVE_GRAPH_SCORE_DECAY: float = 0.85
    RETRIEVE_COMMIT_SCORE_BOOST: float = 1.25

    # Defaults for the simple POST /api/v1/search endpoint
    SEARCH_DEFAULT_REPO_ID: str = ""
    SEARCH_DEFAULT_SNAPSHOT_ID: str = ""

    # Self-referencing URL used by slack_sync background tasks to call ingest/pipeline routes
    EMBEDDING_ENGINE_SELF_URL: str = "http://localhost:8000"

    # Internal maintenance endpoints (DELETE /internal/v1/repos/{repo_id}/vectors)
    GATEWAY_INTERNAL_SERVICE_SECRET: str = ""

    # Redis (streams + general client)
    REDIS_URL: str = "redis://localhost:6379/0"
    REDIS_STREAM_PREFIX: str = ""
    CONSUMER_GROUP: str = "embedding-engine-service"
    CONSUMER_NAME: str = ""
    STREAM_CONSUMER_BLOCK_MS: int = 2000
    STREAM_CONSUMER_BATCH_SIZE: int = 10
    STREAM_MAX_DELIVERY_ATTEMPTS: int = 5
    STREAM_DLQ_SUFFIX: str = "dlq"

    # Celery
    CELERY_BROKER_URL: str = "redis://localhost:6379/0"
    CELERY_RESULT_BACKEND: str = "redis://localhost:6379/1"
    CELERY_WORKER_CONCURRENCY: int = 4
    # Hard/soft caps on a single snapshot-embedding task. A large snapshot can
    # legitimately embed longer than the old 30-min default; too low a limit
    # kills the task mid-run (and, before incremental flushing, lost everything).
    # Soft limit fires first (raises SoftTimeLimitExceeded); hard limit force-kills.
    CELERY_TASK_TIME_LIMIT_SEC: int = 90 * 60
    CELERY_TASK_SOFT_TIME_LIMIT_SEC: int = 85 * 60

    # Tokenizer / retries
    TOKENIZER_MODEL: str = "cl100k_base"
    MAX_TOKENS_PER_CHUNK: int = 8000
    RETRY_MAX_ATTEMPTS: int = 3
    RETRY_BACKOFF_FACTOR: float = 2.0

    @field_validator("EMBEDDING_DIMENSION", "MAX_TOKENS_PER_CHUNK", "RETRY_MAX_ATTEMPTS")
    @classmethod
    def _positive_int(cls, value: int) -> int:
        if value < 1:
            raise ValueError("must be >= 1")
        return value

    @field_validator(
        "EMBEDDING_BATCH_SIZE",
        "CELERY_WORKER_CONCURRENCY",
        "OPENAI_REQUEST_TIMEOUT_SEC",
        "EMBEDDING_MAX_TOKENS_PER_REQUEST",
        "EMBEDDING_UPSERT_FLUSH_SIZE",
        "CELERY_TASK_TIME_LIMIT_SEC",
        "CELERY_TASK_SOFT_TIME_LIMIT_SEC",
    )
    @classmethod
    def _positive_batch(cls, value: int) -> int:
        if value < 1:
            raise ValueError("must be >= 1")
        return value

    @field_validator("OPENAI_EMBEDDING_BATCH_DELAY_MS")
    @classmethod
    def _non_negative_delay(cls, value: int) -> int:
        if value < 0:
            raise ValueError("must be >= 0")
        return value

    def resolved_environment(self) -> str:
        return (self.ENVIRONMENT or "development").strip().lower()

    def is_local(self) -> bool:
        return self.resolved_environment() in LOCAL_ENVIRONMENTS

    def resolved_docs_database(self) -> str:
        name = self.DOCS_MONGODB_DATABASE.strip()
        return name or self.DATABASE_NAME

    def resolved_code_chunks_database(self) -> str:
        name = self.CODE_CHUNKS_MONGODB_DATABASE.strip()
        return name or self.DATABASE_NAME

    def resolved_commit_analyses_database(self) -> str:
        name = self.COMMIT_ANALYSES_MONGODB_DATABASE.strip()
        return name or self.DATABASE_NAME

    def stream_name(self, base: str) -> str:
        """Return a Redis stream name with optional REDIS_STREAM_PREFIX."""
        base = base.strip()
        prefix = self.REDIS_STREAM_PREFIX.strip()
        if not prefix:
            return base
        if not prefix.endswith("."):
            prefix = f"{prefix}."
        return f"{prefix}{base}"

    def validate_runtime(self) -> None:
        """Fail fast when required settings are missing outside local/dev."""
        if self.EMBEDDING_PROVIDER not in SUPPORTED_EMBEDDING_PROVIDERS:
            raise ValueError(
                f"Unsupported EMBEDDING_PROVIDER: {self.EMBEDDING_PROVIDER}. "
                f"Supported: {', '.join(sorted(SUPPORTED_EMBEDDING_PROVIDERS))}"
            )

        if self.is_local():
            return

        missing: list[str] = []
        for name, value in (
            ("MONGO_URI", self.MONGO_URI),
            ("OPENAI_API_KEY", self.OPENAI_API_KEY),
            ("QDRANT_URL", self.QDRANT_URL),
            ("QDRANT_COLLECTION_NAME", self.QDRANT_COLLECTION_NAME),
            ("REDIS_URL", self.REDIS_URL),
            ("CELERY_BROKER_URL", self.CELERY_BROKER_URL),
            ("CELERY_RESULT_BACKEND", self.CELERY_RESULT_BACKEND),
        ):
            if not str(value).strip():
                missing.append(name)

        if missing:
            raise ValueError(
                f"Missing required configuration: {', '.join(missing)}"
            )


settings = Settings()
