Skip to content
22 changes: 13 additions & 9 deletions src/bandai/config/__init__.py
Original file line number Diff line number Diff line change
@@ -1,18 +1,22 @@
from .config import (
from ._constants import ( # noqa: F401
_MAX_REVIEW_ITERATIONS,
EnvOverrides,
PROVIDERS,
LLMProfile,
ProviderProfile,
clear_caches,
)
from ._env import ( # noqa: F401
EnvOverrides,
get_active_provider,
get_api_key,
)
from .llm import get_llm # noqa: F401
from .embedder import ( # noqa: F401
get_embedder,
get_llm,
get_memory,
validate_config,
PROVIDERS,
ProviderProfile,
LLMProfile,
)
from .portals import (
from .memory import get_memory # noqa: F401
from .validation import validate_config # noqa: F401
from .portals import ( # noqa: F401
BANDI_PORTALS,
PORTAL_WEIGHTS,
PortalConfig,
Expand Down
104 changes: 104 additions & 0 deletions src/bandai/config/_constants.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,104 @@
from __future__ import annotations

import logging
import os
from typing import TYPE_CHECKING

from pydantic import BaseModel, Field

if TYPE_CHECKING:
pass

log = logging.getLogger(__name__)
logging.getLogger("opentelemetry.attributes").setLevel(logging.ERROR)


# Provider Profile Models


class LLMProfile(BaseModel):
"""Configuration for a single LLM model slot (main or fast)."""

model: str
temperature: float = Field(default=0.3, ge=0.0, le=2.0)
max_tokens: int = Field(default=4096, ge=1, le=131072)


class ProviderProfile(BaseModel):
"""Configuration for a complete LLM provider."""

name: str
description: str = ""
base_url: str
main: LLMProfile
fast: LLMProfile
env_key: str = "API_KEY"

@property
def main_model(self) -> str:
"""Full model identifier in provider/model format."""
if "/" in self.main.model:
return self.main.model
return f"{self.name}/{self.main.model}"

@property
def fast_model(self) -> str:
"""Full model identifier in provider/model format."""
if "/" in self.fast.model:
return self.fast.model
return f"{self.name}/{self.fast.model}"


# Built-in Provider Profiles

PROVIDERS: dict[str, ProviderProfile] = {
"openrouter": ProviderProfile(
name="openrouter",
description="OpenRouter - multi-provider routing (supports OpenAI, Anthropic, Google, etc.)",
base_url="https://openrouter.ai/api/v1",
main=LLMProfile(model="anthropic/claude-sonnet-4-20250514", temperature=0.3, max_tokens=4096),
fast=LLMProfile(model="openai/gpt-4o-mini", temperature=0.3, max_tokens=4096),
env_key="OPENROUTER_API_KEY",
),
"anthropic": ProviderProfile(
name="anthropic",
description="Anthropic - direct API access",
base_url="https://api.anthropic.com/v1",
main=LLMProfile(model="claude-sonnet-4-20250514", temperature=0.3, max_tokens=4096),
fast=LLMProfile(model="claude-haiku-3-5-20241022", temperature=0.3, max_tokens=4096),
env_key="ANTHROPIC_API_KEY",
),
"openai": ProviderProfile(
name="openai",
description="OpenAI - direct API access",
base_url="https://api.openai.com/v1",
main=LLMProfile(model="gpt-4o", temperature=0.3, max_tokens=4096),
fast=LLMProfile(model="gpt-4o-mini", temperature=0.3, max_tokens=4096),
env_key="OPENAI_API_KEY",
),
"ollama": ProviderProfile(
name="ollama",
description="Ollama - local LLM server",
base_url="http://localhost:11434/v1",
main=LLMProfile(model="llama3", temperature=0.3, max_tokens=4096),
fast=LLMProfile(model="llama3", temperature=0.3, max_tokens=4096),
env_key="",
),
}


# Review Loop Limits

_MAX_REVIEW_ITERATIONS = int(os.getenv("MAX_REVIEW_ITERATIONS", "5"))


# Cache Management (testing)


def clear_caches() -> None:
"""Clear all LLM/embedder caches. Useful in tests."""
from bandai.config.llm import get_llm # noqa: E402
from bandai.config.embedder import get_embedder # noqa: E402

get_llm.cache_clear()
get_embedder.cache_clear()
81 changes: 81 additions & 0 deletions src/bandai/config/_env.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,81 @@
from __future__ import annotations

import os

from pydantic import BaseModel, Field

from bandai.config._constants import PROVIDERS, ProviderProfile # noqa: E402

# Environment Override Model


class EnvOverrides(BaseModel):
"""Parsed LLM environment variable overrides."""

model: str | None = None
temperature: float | None = Field(default=None, ge=0.0, le=2.0)
max_tokens: int | None = Field(default=None, ge=1, le=131072)

@classmethod
def from_env(cls, prefix: str) -> EnvOverrides:
"""Read overrides for *prefix* (MAIN_ or FAST_)."""

def _float(key: str) -> float | None:
raw = os.getenv(key)
if raw is None:
return None
try:
return float(raw)
except ValueError:
return None

def _int(key: str) -> int | None:
raw = os.getenv(key)
if raw is None:
return None
try:
return int(raw)
except ValueError:
return None

return cls(
model=os.getenv(f"{prefix}MODEL") or None,
temperature=_float(f"{prefix}TEMPERATURE"),
max_tokens=_int(f"{prefix}MAX_TOKENS"),
)


# Active Configuration


def get_active_provider() -> ProviderProfile:
"""Return the active provider profile based on LLM_PROVIDER env var."""
provider_name = os.getenv("LLM_PROVIDER", "openrouter")
if provider_name not in PROVIDERS:
available = ", ".join(sorted(PROVIDERS.keys()))
raise ValueError(f"Unknown LLM_PROVIDER '{provider_name}'. " f"Available providers: {available}")
return PROVIDERS[provider_name]


def get_api_key() -> str:
"""Return the API key for the active provider."""
provider = get_active_provider()

# Provider-specific key takes priority
if provider.env_key:
key = os.getenv(provider.env_key, "")
if key and key != f"your_{provider.env_key.lower()}_here":
return key

# Generic fallback
key = os.getenv("API_KEY", "")
if key and key != "your_api_key_here":
return key

available_msg = ""
if provider.env_key:
available_msg = f" Set {provider.env_key} or API_KEY in your .env file."
else:
available_msg = " No API key required for this provider."

raise ValueError(f"No API key configured for provider '{provider.name}'.{available_msg}")
Loading
Loading