Skip to content

forge.ai

Unified AI model interface across OpenAI, Anthropic, Gemini, and Ollama.

forge.ai

AI model abstraction module — unified interface across LLM providers.

Provides a single API surface for completions and streaming across OpenAI, Anthropic, and a fully offline MockAdapter for testing. Includes token estimation, cost calculation, fallback routing, and pre-request budget checking.

Classes

AIError

Bases: ForgeError

Base exception for all AI module errors.

Source code in src/forge/ai/exceptions.py
class AIError(ForgeError):
    """Base exception for all AI module errors."""

AIModule

Bases: ForgeModule

Source code in src/forge/ai/module.py
class AIModule(ForgeModule):
    name = "ai"
    dependencies: ClassVar[list[str]] = ["config", "retry"]

    def __init__(self) -> None:
        super().__init__()
        self._router: ModelRouter | None = None
        self._config_ai: Any = None

        # Metrics
        self._request_count = 0
        self._latency_ms_history: list[float] = []
        self._input_tokens = 0
        self._output_tokens = 0
        self._total_cost = 0.0

        # OTel instruments (lazily created)
        self._otel_request_counter: Any = None
        self._otel_latency_histogram: Any = None
        self._otel_token_counter: Any = None
        self._otel_cost_counter: Any = None

    def _ensure_otel_instruments(self) -> None:
        if self._otel_request_counter is not None:
            return
        meter = get_meter()
        if meter is None:
            return
        self._otel_request_counter = meter.create_counter(
            "ai.request.count",
            description="Total number of AI completion requests",
            unit="1",
        )
        self._otel_latency_histogram = meter.create_histogram(
            "ai.latency",
            description="Latency of AI completion requests in milliseconds",
            unit="ms",
        )
        self._otel_token_counter = meter.create_counter(
            "ai.token.count",
            description="Total number of tokens used across AI requests",
            unit="1",
        )
        self._otel_cost_counter = meter.create_counter(
            "ai.cost.estimate",
            description="Estimated cost of AI requests in USD",
            unit="USD",
        )

    # ── Public API ─────────────────────────────────────────────────

    @property
    def router(self) -> ModelRouter:
        if self._router is None:
            raise RuntimeError("AIModule not initialised")
        return self._router

    async def complete(
        self,
        request: CompletionRequest,
        output_schema: type[BaseModel] | None = None,
        fallback_models: list[str] | None = None,
        max_retries: int | None = None,
    ) -> CompletionResponse | BaseModel:
        import time

        start_time = time.monotonic()

        if output_schema is not None:
            from forge.ai.structured import StructuredOutputEnforcer

            retries = (
                max_retries
                if max_retries is not None
                else (self._config_ai.structured_output_retries if self._config_ai else 3)
            )
            enforcer = StructuredOutputEnforcer(max_retries=retries)

            async def _complete_fn(msgs: list[Message]) -> CompletionResponse:
                req = request.model_copy(update={"messages": msgs})
                res = await self.router.complete(req, fallback_models)

                if res.usage is not None:
                    try:
                        adapter = self.router.resolve(res.model)
                        if adapter is not None:
                            res.cost = adapter.estimate_cost(res.usage, res.model)
                    except Exception:
                        from forge.ai.tokens import TokenCounter

                        res.cost = TokenCounter.estimate_cost(res.usage, res.model)

                    self._input_tokens += res.usage.input_tokens
                    self._output_tokens += res.usage.output_tokens
                    self._total_cost += res.cost
                return res

            try:
                prev_input = self._input_tokens
                prev_output = self._output_tokens
                prev_cost = self._total_cost

                validated = await enforcer.enforce(request.messages, output_schema, _complete_fn)
                latency_ms = (time.monotonic() - start_time) * 1000.0
                self._request_count += 1
                self._latency_ms_history.append(latency_ms)
                self._record_otel_metrics(
                    latency_ms,
                    input_delta=self._input_tokens - prev_input,
                    output_delta=self._output_tokens - prev_output,
                    cost_delta=self._total_cost - prev_cost,
                )
                return validated
            except Exception:
                latency_ms = (time.monotonic() - start_time) * 1000.0
                self._request_count += 1
                self._latency_ms_history.append(latency_ms)
                self._record_otel_metrics(latency_ms)
                raise
        else:
            try:
                response = await self.router.complete(request, fallback_models)
                latency_ms = (time.monotonic() - start_time) * 1000.0
                response.latency_ms = latency_ms

                input_delta = 0
                output_delta = 0
                cost_delta = 0.0

                if response.usage is not None:
                    try:
                        adapter = self.router.resolve(response.model)
                        if adapter is not None:
                            response.cost = adapter.estimate_cost(response.usage, response.model)
                    except Exception:
                        from forge.ai.tokens import TokenCounter

                        response.cost = TokenCounter.estimate_cost(response.usage, response.model)

                    input_delta = response.usage.input_tokens
                    output_delta = response.usage.output_tokens
                    cost_delta = response.cost
                    self._input_tokens += input_delta
                    self._output_tokens += output_delta
                    self._total_cost += cost_delta

                # Update metrics
                self._request_count += 1
                self._latency_ms_history.append(latency_ms)
                self._record_otel_metrics(
                    latency_ms,
                    input_delta=input_delta,
                    output_delta=output_delta,
                    cost_delta=cost_delta,
                )
                return response
            except Exception:
                latency_ms = (time.monotonic() - start_time) * 1000.0
                self._request_count += 1
                self._latency_ms_history.append(latency_ms)
                self._record_otel_metrics(latency_ms)
                raise

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        import time

        from forge.ai.tokens import TokenCounter

        start_time = time.monotonic()
        input_tokens = TokenCounter.count_messages(request.messages)
        output_tokens = 0

        model_used = request.model

        try:
            async for chunk in self.router.stream(request):
                if chunk.delta:
                    output_tokens += TokenCounter.count_tokens(chunk.delta)
                if chunk.model:
                    model_used = chunk.model
                if chunk.usage and chunk.usage.output_tokens > 0:
                    output_tokens = chunk.usage.output_tokens
                yield chunk

            latency_ms = (time.monotonic() - start_time) * 1000.0

            from forge.ai.models import Usage

            usage = Usage(input_tokens=input_tokens, output_tokens=output_tokens)
            cost = 0.0
            try:
                adapter = self.router.resolve(model_used)
                if adapter is not None:
                    cost = adapter.estimate_cost(usage, model_used)
            except Exception:
                cost = TokenCounter.estimate_cost(usage, model_used)

            self._request_count += 1
            self._latency_ms_history.append(latency_ms)
            self._input_tokens += input_tokens
            self._output_tokens += output_tokens
            self._total_cost += cost
            self._record_otel_metrics(
                latency_ms, input_delta=input_tokens, output_delta=output_tokens, cost_delta=cost
            )

        except Exception:
            latency_ms = (time.monotonic() - start_time) * 1000.0
            self._request_count += 1
            self._latency_ms_history.append(latency_ms)
            self._record_otel_metrics(latency_ms)
            raise

    def _record_otel_metrics(
        self,
        latency_ms: float,
        input_delta: int = 0,
        output_delta: int = 0,
        cost_delta: float = 0.0,
    ) -> None:
        self._ensure_otel_instruments()
        if self._otel_request_counter is not None:
            self._otel_request_counter.add(
                1, {"model": str(self._config_ai.default_model if self._config_ai else "unknown")}
            )
        if self._otel_latency_histogram is not None:
            self._otel_latency_histogram.record(latency_ms)
        if self._otel_token_counter is not None and (input_delta or output_delta):
            self._otel_token_counter.add(input_delta + output_delta, {"type": "total"})
        if self._otel_cost_counter is not None and cost_delta:
            self._otel_cost_counter.add(cost_delta)

    def get_metrics(self) -> dict[str, Any]:
        """Return the collected observability metrics."""
        return {
            "request_count": self._request_count,
            "total_latency_ms": sum(self._latency_ms_history),
            "latency_history": list(self._latency_ms_history),
            "token_count": {
                "input": self._input_tokens,
                "output": self._output_tokens,
                "total": self._input_tokens + self._output_tokens,
            },
            "cost_estimate": round(self._total_cost, 6),
        }

    # ── ForgeModule ────────────────────────────────────────────────

    async def setup(self, runtime: ForgeRuntime) -> None:
        from forge.config.module import ConfigModule

        config_module: ConfigModule = runtime.get(ConfigModule)  # type: ignore[assignment]
        self._config_ai = config_module.config.ai

        self._router = ModelRouter(
            max_tokens_limit=self._config_ai.max_tokens,
            retry_module=runtime.get(RetryModule),
        )

        # Register adapters
        openai_adapter = OpenAIAdapter(
            api_key=self._get_openai_key(),
            timeout=self._config_ai.timeout,
        )
        self._router.register("gpt-4o*", openai_adapter)
        self._router.register("gpt-4*", openai_adapter)
        self._router.register("gpt-3.5*", openai_adapter)

        anthropic_adapter = AnthropicAdapter(
            api_key=self._get_anthropic_key(),
            timeout=self._config_ai.timeout,
        )
        self._router.register("claude*", anthropic_adapter)

        gemini_adapter = GeminiAdapter(
            api_key=self._get_gemini_key(),
            timeout=self._config_ai.timeout,
        )
        self._router.register("gemini*", gemini_adapter)

        ollama_adapter = OllamaAdapter(
            timeout=self._config_ai.timeout,
        )
        self._router.register("ollama*", ollama_adapter)

        mock_adapter = MockAdapter()
        self._router.register("*", mock_adapter)

        # Fallback chain (tried in order)
        for fb_model in self._config_ai.fallback_models:
            adapter_for_fallback = self._router.resolve(fb_model)
            if adapter_for_fallback is not None:
                self._router.register(fb_model, adapter_for_fallback, is_fallback=True)

    async def teardown(self) -> None:
        self._router = None
        self._config_ai = None

    def health_check(self) -> HealthResult:
        if self._router is None:
            return HealthResult.error("AI module not initialised")

        default_model = self._config_ai.default_model
        try:
            adapter = self._router.resolve(default_model)
            if adapter is None:
                return HealthResult.error(
                    f"No adapter registered for default model '{default_model}'"
                )

            if adapter.provider == "mock":
                return HealthResult.ok()

            from forge.ai.models import CompletionRequest, Message
            from forge.core.async_bridge import run_async_health_check

            async def _ping() -> None:
                req = CompletionRequest(
                    model=default_model,
                    messages=[Message(role="user", content="ping")],
                    max_tokens=1,
                )
                await asyncio.wait_for(adapter.complete(req), timeout=5.0)

            run_async_health_check(_ping())
            return HealthResult.ok()
        except Exception as e:
            provider_name = adapter.provider if "adapter" in locals() and adapter else "unknown"
            return HealthResult.error(f"Health check failed for provider '{provider_name}': {e}")

    # ── Internal ───────────────────────────────────────────────────

    def _get_openai_key(self) -> str | None:
        if self._config_ai is None:
            return None
        val = self._config_ai.openai_api_key
        return val.get_secret_value() if val is not None else None

    def _get_anthropic_key(self) -> str | None:
        if self._config_ai is None:
            return None
        val = self._config_ai.anthropic_api_key
        return val.get_secret_value() if val is not None else None

    def _get_gemini_key(self) -> str | None:
        if self._config_ai is None:
            return None
        val = getattr(self._config_ai, "gemini_api_key", None)
        return val.get_secret_value() if val is not None else None
Methods:
get_metrics
get_metrics() -> dict[str, Any]

Return the collected observability metrics.

Source code in src/forge/ai/module.py
def get_metrics(self) -> dict[str, Any]:
    """Return the collected observability metrics."""
    return {
        "request_count": self._request_count,
        "total_latency_ms": sum(self._latency_ms_history),
        "latency_history": list(self._latency_ms_history),
        "token_count": {
            "input": self._input_tokens,
            "output": self._output_tokens,
            "total": self._input_tokens + self._output_tokens,
        },
        "cost_estimate": round(self._total_cost, 6),
    }

AIProviderError

Bases: AIError

Raised when the provider API returns an error (auth, billing, bad request).

Source code in src/forge/ai/exceptions.py
class AIProviderError(AIError):
    """Raised when the provider API returns an error (auth, billing, bad request)."""

AllModelsFailedError

Bases: AIError

Raised when all models in the fallback chain have failed to execute.

Source code in src/forge/ai/exceptions.py
class AllModelsFailedError(AIError):
    """Raised when all models in the fallback chain have failed to execute."""

AnthropicAdapter

Bases: BaseAdapter

Thin wrapper around the Anthropic Python SDK.

Requires anthropic to be installed. Falls back to mock behaviour when the library is unavailable.

Source code in src/forge/ai/adapters/anthropic.py
class AnthropicAdapter(BaseAdapter):
    """
    Thin wrapper around the Anthropic Python SDK.

    Requires ``anthropic`` to be installed.  Falls back to mock behaviour
    when the library is unavailable.
    """

    def __init__(self, api_key: str | None = None, timeout: float = 30.0) -> None:
        self._api_key = api_key
        self._timeout = timeout
        self._client: Any = None
        self._init_client()

    @property
    def provider(self) -> str:
        return "anthropic"

    # ── BaseAdapter ────────────────────────────────────────────────

    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        if self._client is None or not self._api_key:
            return _mock_complete(request)

        payload = self._build_payload(request)
        try:
            response = await self._client.messages.create(**payload)
        except Exception as exc:
            import anthropic

            if isinstance(exc, anthropic.RateLimitError):
                raise RateLimitError(f"Anthropic API rate limit exceeded: {exc}") from exc
            if isinstance(exc, anthropic.AuthenticationError):
                raise AIProviderError(f"Anthropic API authentication failed: {exc}") from exc
            if isinstance(exc, anthropic.NotFoundError):
                raise ModelNotFoundError(f"Anthropic model not found: {exc}") from exc
            raise AIProviderError(f"Anthropic API error: {exc}") from exc

        return CompletionResponse(
            model=response.model or request.model,
            message=Message(
                role="assistant",
                content=response.content[0].text if response.content else "",
            ),
            usage=Usage(
                input_tokens=response.usage.input_tokens if response.usage else 0,
                output_tokens=response.usage.output_tokens if response.usage else 0,
            ),
            provider="anthropic",
        )

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        if self._client is None or not self._api_key:
            async for chunk in _mock_stream(request):
                yield chunk
            return

        payload = self._build_payload(request)
        payload["stream"] = True

        input_tokens = 0
        output_tokens = 0

        try:
            async with self._client.messages.create(**payload) as msg_stream:
                async for event in msg_stream:
                    delta_content = ""
                    finish_reason = None
                    usage = None

                    if event.type == "message_start":
                        if hasattr(event.message, "usage") and event.message.usage:
                            input_tokens = event.message.usage.input_tokens
                    elif event.type == "content_block_delta":
                        if hasattr(event.delta, "text") and event.delta.text:
                            delta_content = event.delta.text
                    elif event.type == "message_delta":
                        if hasattr(event, "usage") and event.usage:
                            output_tokens = event.usage.output_tokens
                        if hasattr(event.delta, "stop_reason") and event.delta.stop_reason:
                            finish_reason = event.delta.stop_reason
                    elif event.type == "message_stop":
                        usage = Usage(
                            input_tokens=input_tokens,
                            output_tokens=output_tokens,
                        )

                    if delta_content or finish_reason or usage:
                        yield StreamChunk(
                            delta=delta_content,
                            finish_reason=finish_reason,
                            usage=usage,
                            model=request.model,
                            provider="anthropic",
                        )
        except Exception as exc:
            import anthropic

            if isinstance(exc, anthropic.RateLimitError):
                raise RateLimitError(f"Anthropic API rate limit exceeded: {exc}") from exc
            if isinstance(exc, anthropic.AuthenticationError):
                raise AIProviderError(f"Anthropic API authentication failed: {exc}") from exc
            if isinstance(exc, anthropic.NotFoundError):
                raise ModelNotFoundError(f"Anthropic model not found: {exc}") from exc
            raise StreamInterruptedError(f"Anthropic stream interrupted: {exc}") from exc

    def count_tokens(self, text: str) -> int:
        return TokenCounter.count_tokens(text)

    # ── Internal ───────────────────────────────────────────────────

    def _init_client(self) -> None:
        try:
            import anthropic

            if not self._api_key:
                self._client = None
            else:
                self._client = anthropic.AsyncAnthropic(
                    api_key=self._api_key, timeout=self._timeout
                )
        except (ImportError, Exception) as exc:
            _logger.warning("Failed to initialize Anthropic client: %s — using mock", exc)
            self._client = None

    def _build_payload(self, request: CompletionRequest) -> dict[str, Any]:
        payload: dict[str, Any] = {
            "model": request.model,
            "max_tokens": request.max_tokens or 1024,
            "messages": [
                {"role": m.role, "content": m.content}
                for m in request.messages
                if m.role != "system"
            ],
        }
        system_msgs = [m.content for m in request.messages if m.role == "system"]
        if system_msgs:
            payload["system"] = "\n".join(system_msgs)
        if request.temperature is not None:
            payload["temperature"] = request.temperature
        payload.update(request.extra)
        return payload

BaseAdapter

Bases: ABC

Interface that every provider adapter must implement.

Source code in src/forge/ai/adapters/base.py
class BaseAdapter(ABC):
    """Interface that every provider adapter must implement."""

    @property
    @abstractmethod
    def provider(self) -> str:
        """Short label such as ``"openai"`` or ``"anthropic"``."""

    @abstractmethod
    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        """Synchronous (non-streaming) completion."""

    async def stream(self, _request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        """Yield content chunks as they arrive."""
        # Subclasses override; the ``if False: yield`` makes type
        # checkers recognise this as an async generator.
        if False:  # pragma: no cover
            yield
        raise NotImplementedError

    @abstractmethod
    def count_tokens(self, text: str) -> int:
        """Return an estimated token count for *text*."""

    def estimate_cost(self, usage: Usage, model: str) -> float:
        """Return the estimated USD cost for the given usage."""
        return TokenCounter.estimate_cost(usage, model)
Attributes
provider abstractmethod property
provider: str

Short label such as "openai" or "anthropic".

Methods:
complete abstractmethod async
complete(request: CompletionRequest) -> CompletionResponse

Synchronous (non-streaming) completion.

Source code in src/forge/ai/adapters/base.py
@abstractmethod
async def complete(self, request: CompletionRequest) -> CompletionResponse:
    """Synchronous (non-streaming) completion."""
count_tokens abstractmethod
count_tokens(text: str) -> int

Return an estimated token count for text.

Source code in src/forge/ai/adapters/base.py
@abstractmethod
def count_tokens(self, text: str) -> int:
    """Return an estimated token count for *text*."""
estimate_cost
estimate_cost(usage: Usage, model: str) -> float

Return the estimated USD cost for the given usage.

Source code in src/forge/ai/adapters/base.py
def estimate_cost(self, usage: Usage, model: str) -> float:
    """Return the estimated USD cost for the given usage."""
    return TokenCounter.estimate_cost(usage, model)
stream async
stream(_request: CompletionRequest) -> AsyncIterator[StreamChunk]

Yield content chunks as they arrive.

Source code in src/forge/ai/adapters/base.py
async def stream(self, _request: CompletionRequest) -> AsyncIterator[StreamChunk]:
    """Yield content chunks as they arrive."""
    # Subclasses override; the ``if False: yield`` makes type
    # checkers recognise this as an async generator.
    if False:  # pragma: no cover
        yield
    raise NotImplementedError

BudgetExceededError

Bases: TokenLimitError

Raised when the estimated token count exceeds the configured limit.

Source code in src/forge/ai/tokens.py
class BudgetExceededError(TokenLimitError):
    """Raised when the estimated token count exceeds the configured limit."""

GeminiAdapter

Bases: BaseAdapter

Adapter for Google Gemini models via google-generativeai.

Source code in src/forge/ai/adapters/gemini.py
class GeminiAdapter(BaseAdapter):
    """Adapter for Google Gemini models via google-generativeai."""

    def __init__(self, api_key: str | None = None, timeout: float = 30.0) -> None:
        self._api_key = api_key
        self._timeout = timeout
        self._client: Any = None
        self._init_client()

    @property
    def provider(self) -> str:
        return "gemini"

    def _init_client(self) -> None:
        try:
            import google.generativeai as genai

            if self._api_key:
                genai.configure(api_key=self._api_key)  # type: ignore[attr-defined,unused-ignore]
            self._client = genai
        except ImportError:
            _logger.warning("google-generativeai package not installed — using mock fallback")
            self._client = None

    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        if self._client is None or not self._api_key:
            return self._mock_complete(request)

        # Prepare messages in Gemini format
        system_instruction = None
        gemini_messages = []
        for msg in request.messages:
            if msg.role == "system":
                if system_instruction is None:
                    system_instruction = msg.content
                else:
                    system_instruction += "\n" + msg.content
            elif msg.role == "user":
                gemini_messages.append({"role": "user", "parts": [msg.content]})
            elif msg.role in ("assistant", "model"):
                gemini_messages.append({"role": "model", "parts": [msg.content]})

        # Initialize the generative model
        try:
            model = self._client.GenerativeModel(
                model_name=request.model,
                system_instruction=system_instruction,
            )
        except Exception as exc:
            raise AIProviderError(
                f"Failed to initialize Gemini model '{request.model}': {exc}"
            ) from exc

        # Set up generation config
        generation_config: dict[str, Any] = {}
        if request.max_tokens is not None:
            generation_config["max_output_tokens"] = request.max_tokens
        if request.temperature is not None:
            generation_config["temperature"] = request.temperature

        # Call the Gemini API
        try:
            response = await model.generate_content_async(
                contents=gemini_messages,
                generation_config=self._client.types.GenerationConfig(**generation_config),
                request_options={"timeout": self._timeout},
            )
        except Exception as exc:
            _raise_gemini_error(exc)

        # Parse output and usage
        content = response.text or ""
        prompt_tokens = 0
        completion_tokens = 0
        if response.usage_metadata:
            prompt_tokens = response.usage_metadata.prompt_token_count
            completion_tokens = response.usage_metadata.candidates_token_count

        return CompletionResponse(
            model=request.model,
            message=Message(role="assistant", content=content),
            usage=Usage(input_tokens=prompt_tokens, output_tokens=completion_tokens),
            provider="gemini",
        )

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        if self._client is None or not self._api_key:
            async for chunk in _mock_stream(request):
                yield chunk
            return

        # Prepare messages in Gemini format
        system_instruction = None
        gemini_messages = []
        for msg in request.messages:
            if msg.role == "system":
                if system_instruction is None:
                    system_instruction = msg.content
                else:
                    system_instruction += "\n" + msg.content
            elif msg.role == "user":
                gemini_messages.append({"role": "user", "parts": [msg.content]})
            elif msg.role in ("assistant", "model"):
                gemini_messages.append({"role": "model", "parts": [msg.content]})

        # Initialize the generative model
        try:
            model = self._client.GenerativeModel(
                model_name=request.model,
                system_instruction=system_instruction,
            )
        except Exception as exc:
            raise AIProviderError(
                f"Failed to initialize Gemini model '{request.model}': {exc}"
            ) from exc

        # Set up generation config
        generation_config: dict[str, Any] = {}
        if request.max_tokens is not None:
            generation_config["max_output_tokens"] = request.max_tokens
        if request.temperature is not None:
            generation_config["temperature"] = request.temperature

        # Initiate the streaming request
        try:
            stream_resp = await model.generate_content_async(
                contents=gemini_messages,
                generation_config=self._client.types.GenerationConfig(**generation_config),
                stream=True,
                request_options={"timeout": self._timeout},
            )
        except Exception as exc:
            _raise_gemini_error(
                exc, "starting Gemini stream", default_factory=StreamInterruptedError
            )

        try:
            async for chunk in stream_resp:
                delta_content = ""
                finish_reason = None
                usage = None

                # Check for content blocked by safety filters
                if hasattr(chunk, "prompt_feedback") and chunk.prompt_feedback:
                    block_reason = chunk.prompt_feedback.block_reason
                    if block_reason:
                        yield StreamChunk(
                            delta="",
                            finish_reason="content_filter",
                            usage=None,
                            model=request.model,
                            provider="gemini",
                        )
                        return

                # Extract text from candidates
                if chunk.candidates:
                    candidate = chunk.candidates[0]

                    if candidate.content and candidate.content.parts:
                        parts_text = "".join(
                            p.text for p in candidate.content.parts if hasattr(p, "text") and p.text
                        )
                        if parts_text:
                            delta_content = parts_text

                    if candidate.finish_reason:
                        finish_reason = _map_gemini_finish_reason(candidate.finish_reason)

                # Extract usage metadata (typically available on the last chunk)
                if chunk.usage_metadata:
                    usage = Usage(
                        input_tokens=getattr(chunk.usage_metadata, "prompt_token_count", 0),
                        output_tokens=getattr(chunk.usage_metadata, "candidates_token_count", 0),
                    )

                if delta_content or finish_reason or usage:
                    yield StreamChunk(
                        delta=delta_content,
                        finish_reason=finish_reason,
                        usage=usage,
                        model=request.model,
                        provider="gemini",
                    )
        except Exception as exc:
            _raise_gemini_error(exc, "Gemini stream", default_factory=StreamInterruptedError)

    def count_tokens(self, text: str) -> int:
        return TokenCounter.count_tokens(text)

    def _mock_complete(self, request: CompletionRequest) -> CompletionResponse:
        return CompletionResponse(
            model=request.model,
            message=Message(role="assistant", content="[mock gemini response]"),
            usage=Usage(
                input_tokens=TokenCounter.count_messages(request.messages),
                output_tokens=10,
            ),
            provider="gemini",
        )

MockAdapter

Bases: BaseAdapter

Fully offline adapter for testing — no external API calls.

Returns a canned response and never raises rate-limit or auth errors.

Source code in src/forge/ai/adapters/mock.py
class MockAdapter(BaseAdapter):
    """
    Fully offline adapter for testing — no external API calls.

    Returns a canned response and never raises rate-limit or auth errors.
    """

    def __init__(self, response_text: str = "[mock response]") -> None:
        self._response_text = response_text

    @property
    def provider(self) -> str:
        return "mock"

    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        return CompletionResponse(
            model=request.model,
            message=Message(role="assistant", content=self._response_text),
            usage=Usage(
                input_tokens=TokenCounter.count_messages(request.messages),
                output_tokens=self.count_tokens(self._response_text),
            ),
            provider="mock",
        )

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        yield StreamChunk(
            delta=self._response_text,
            finish_reason="stop",
            usage=Usage(
                input_tokens=TokenCounter.count_messages(request.messages),
                output_tokens=self.count_tokens(self._response_text),
            ),
            model=request.model,
            provider="mock",
        )

    def count_tokens(self, text: str) -> int:
        return TokenCounter.count_tokens(text)

ModelNotFoundError

Bases: AIProviderError

Raised when the requested model is not available or registered.

Source code in src/forge/ai/exceptions.py
class ModelNotFoundError(AIProviderError):
    """Raised when the requested model is not available or registered."""

ModelRouter

Selects an adapter based on model name and supports fallback chains.

Usage::

router = ModelRouter()
router.register("gpt-4o*", openai_adapter)
router.register("claude*", anthropic_adapter)
router.register("*", mock_adapter)

resp = await router.complete(request)
Source code in src/forge/ai/router.py
class ModelRouter:
    """
    Selects an adapter based on model name and supports fallback chains.

    Usage::

        router = ModelRouter()
        router.register("gpt-4o*", openai_adapter)
        router.register("claude*", anthropic_adapter)
        router.register("*", mock_adapter)

        resp = await router.complete(request)
    """

    def __init__(
        self,
        max_tokens_limit: int = 128_000,
        retry_module: Any | None = None,
    ) -> None:
        self._patterns: list[tuple[re.Pattern[str], BaseAdapter]] = []
        self._fallback: list[tuple[str, BaseAdapter]] = []
        self._max_tokens_limit = max_tokens_limit
        self._retry_module = retry_module

    # ── Registration ───────────────────────────────────────────────

    def register(
        self,
        model_pattern: str,
        adapter: BaseAdapter,
        *,
        is_fallback: bool = False,
    ) -> None:
        """
        Register *adapter* for models matching *model_pattern* (glob-style).

        Use ``"*"`` as the pattern to match all models.  When
        ``is_fallback=True`` the adapter is added to the fallback chain.
        """
        regex = _glob_to_regex(model_pattern)
        compiled = re.compile(regex)
        if is_fallback:
            self._fallback.append((model_pattern, adapter))
        else:
            self._patterns.append((compiled, adapter))

    # ── Resolution ─────────────────────────────────────────────────

    def resolve(self, model: str) -> BaseAdapter | None:
        """Return the first adapter whose pattern matches *model*, or ``None``."""
        for pattern, adapter in self._patterns:
            if pattern.match(model):
                return adapter
        return None

    # ── Public API ─────────────────────────────────────────────────

    async def complete(
        self,
        request: CompletionRequest,
        fallback_models: list[str] | None = None,
    ) -> CompletionResponse:
        """
        Execute the request with automatic fallback on failure.

        *fallback_models* lists model names to try in order after the
        primary model fails.  If ``None``, uses the router's internal
        fallback adapter list.
        """
        adapter = self.resolve(request.model)
        if adapter is None:
            raise ModelNotFoundError(f"No adapter registered for model {request.model!r}")

        # Pre-request budget check
        TokenCounter.check_budget(request, self._max_tokens_limit)

        # Attempt primary adapter
        last_error: Exception | None = None
        try:
            return await self._call_adapter(adapter, request)
        except Exception as exc:
            last_error = exc
            _logger.warning(
                "Model %r via %s failed: %s — trying fallback",
                request.model,
                adapter.provider,
                exc,
            )

        # Fallback chain
        fallback_adapters = self._build_fallback_chain(fallback_models)
        for fb_model, fb_adapter in fallback_adapters:
            try:
                fb_request = request.model_copy(update={"model": fb_model})
                return await self._call_adapter(fb_adapter, fb_request)
            except Exception as exc:
                last_error = exc
                _logger.warning("Fallback %s also failed: %s", fb_adapter.provider, exc)

        raise AllModelsFailedError(
            f"All {1 + len(fallback_adapters)} adapter(s) failed for model {request.model!r}"
        ) from last_error

    async def stream(
        self,
        request: CompletionRequest,
    ) -> AsyncIterator[StreamChunk]:
        adapter = self.resolve(request.model)
        if adapter is None:
            raise ModelNotFoundError(f"No adapter registered for model {request.model!r}")
        TokenCounter.check_budget(request, self._max_tokens_limit)
        async for chunk in adapter.stream(request):
            yield chunk

    async def _call_adapter(
        self, adapter: BaseAdapter, request: CompletionRequest
    ) -> CompletionResponse:
        if self._retry_module is not None:
            from typing import cast

            from forge.ai.exceptions import RateLimitError

            res = await self._retry_module.retry(
                adapter.complete,
                retryable_exceptions=(RateLimitError,),
            )(request)
            return cast("CompletionResponse", res)
        return await adapter.complete(request)

    # ── Internal ───────────────────────────────────────────────────

    def _build_fallback_chain(
        self,
        fallback_models: list[str] | None,
    ) -> list[tuple[str, BaseAdapter]]:
        if fallback_models:
            adapters: list[tuple[str, BaseAdapter]] = []
            for model in fallback_models:
                a = self.resolve(model)
                if a is not None:
                    adapters.append((model, a))
            return adapters

        chain: list[tuple[str, BaseAdapter]] = []
        for pat, adapter in self._fallback:
            model_name = pat.replace("*", "").replace("?", "")
            chain.append((model_name, adapter))
        return chain
Methods:
complete async
complete(request: CompletionRequest, fallback_models: list[str] | None = None) -> CompletionResponse

Execute the request with automatic fallback on failure.

fallback_models lists model names to try in order after the primary model fails. If None, uses the router's internal fallback adapter list.

Source code in src/forge/ai/router.py
async def complete(
    self,
    request: CompletionRequest,
    fallback_models: list[str] | None = None,
) -> CompletionResponse:
    """
    Execute the request with automatic fallback on failure.

    *fallback_models* lists model names to try in order after the
    primary model fails.  If ``None``, uses the router's internal
    fallback adapter list.
    """
    adapter = self.resolve(request.model)
    if adapter is None:
        raise ModelNotFoundError(f"No adapter registered for model {request.model!r}")

    # Pre-request budget check
    TokenCounter.check_budget(request, self._max_tokens_limit)

    # Attempt primary adapter
    last_error: Exception | None = None
    try:
        return await self._call_adapter(adapter, request)
    except Exception as exc:
        last_error = exc
        _logger.warning(
            "Model %r via %s failed: %s — trying fallback",
            request.model,
            adapter.provider,
            exc,
        )

    # Fallback chain
    fallback_adapters = self._build_fallback_chain(fallback_models)
    for fb_model, fb_adapter in fallback_adapters:
        try:
            fb_request = request.model_copy(update={"model": fb_model})
            return await self._call_adapter(fb_adapter, fb_request)
        except Exception as exc:
            last_error = exc
            _logger.warning("Fallback %s also failed: %s", fb_adapter.provider, exc)

    raise AllModelsFailedError(
        f"All {1 + len(fallback_adapters)} adapter(s) failed for model {request.model!r}"
    ) from last_error
register
register(model_pattern: str, adapter: BaseAdapter, *, is_fallback: bool = False) -> None

Register adapter for models matching model_pattern (glob-style).

Use "*" as the pattern to match all models. When is_fallback=True the adapter is added to the fallback chain.

Source code in src/forge/ai/router.py
def register(
    self,
    model_pattern: str,
    adapter: BaseAdapter,
    *,
    is_fallback: bool = False,
) -> None:
    """
    Register *adapter* for models matching *model_pattern* (glob-style).

    Use ``"*"`` as the pattern to match all models.  When
    ``is_fallback=True`` the adapter is added to the fallback chain.
    """
    regex = _glob_to_regex(model_pattern)
    compiled = re.compile(regex)
    if is_fallback:
        self._fallback.append((model_pattern, adapter))
    else:
        self._patterns.append((compiled, adapter))
resolve
resolve(model: str) -> BaseAdapter | None

Return the first adapter whose pattern matches model, or None.

Source code in src/forge/ai/router.py
def resolve(self, model: str) -> BaseAdapter | None:
    """Return the first adapter whose pattern matches *model*, or ``None``."""
    for pattern, adapter in self._patterns:
        if pattern.match(model):
            return adapter
    return None

OllamaAdapter

Bases: BaseAdapter

Adapter for local Ollama models via the Ollama HTTP API.

Requires httpx to be installed. Falls back to mock behaviour when the library is unavailable.

Source code in src/forge/ai/adapters/ollama.py
class OllamaAdapter(BaseAdapter):
    """
    Adapter for local Ollama models via the Ollama HTTP API.

    Requires ``httpx`` to be installed.  Falls back to mock behaviour
    when the library is unavailable.
    """

    def __init__(
        self,
        base_url: str = "http://localhost:11434",
        timeout: float = 120.0,
        mock: bool = False,
    ) -> None:
        self._base_url = base_url.rstrip("/")
        self._timeout = timeout
        self._mock = mock
        self._client: Any = None
        if not mock:
            self._init_client()

    @property
    def provider(self) -> str:
        return "ollama"

    # ── BaseAdapter ────────────────────────────────────────────────

    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        if self._client is None or self._mock:
            return _mock_complete(request)

        payload = self._build_payload(request)
        try:
            import httpx

            response = await self._client.post(
                f"{self._base_url}/api/chat",
                json=payload,
                timeout=self._timeout,
            )
        except Exception as exc:
            _raise_ollama_error(exc)

        try:
            response.raise_for_status()
        except httpx.HTTPStatusError as exc:
            _raise_ollama_error(exc)

        try:
            data = response.json()
        except Exception as exc:
            raise AIProviderError(f"Ollama returned invalid JSON: {exc}") from exc

        content = data.get("message", {}).get("content", "")
        model = data.get("model", request.model)

        usage = Usage(
            input_tokens=data.get("prompt_eval_count", 0),
            output_tokens=data.get("eval_count", 0),
        )

        return CompletionResponse(
            model=model,
            message=Message(role="assistant", content=content),
            usage=usage,
            provider="ollama",
        )

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        if self._client is None or self._mock:
            async for chunk in _mock_stream(request):
                yield chunk
            return

        payload = self._build_payload(request)
        payload["stream"] = True

        try:
            import httpx

            async with self._client.stream(
                "POST",
                f"{self._base_url}/api/chat",
                json=payload,
                timeout=self._timeout,
            ) as response:
                response.raise_for_status()
                async for line in response.aiter_lines():
                    if not line.strip():
                        continue
                    try:
                        data = json.loads(line)
                    except json.JSONDecodeError as exc:
                        _logger.warning("Skipping malformed Ollama stream line: %s", exc)
                        continue

                    delta_content = data.get("message", {}).get("content", "")
                    done = data.get("done", False)
                    finish_reason = "stop" if done else None

                    usage = None
                    if done and "prompt_eval_count" in data:
                        usage = Usage(
                            input_tokens=data.get("prompt_eval_count", 0),
                            output_tokens=data.get("eval_count", 0),
                        )

                    model = data.get("model", request.model)

                    if delta_content or finish_reason or usage:
                        yield StreamChunk(
                            delta=delta_content,
                            finish_reason=finish_reason,
                            usage=usage,
                            model=model,
                            provider="ollama",
                        )
        except httpx.HTTPStatusError as exc:
            _raise_ollama_error(exc)
        except Exception as exc:
            raise StreamInterruptedError(f"Ollama stream interrupted: {exc}") from exc

    def count_tokens(self, text: str) -> int:
        return TokenCounter.count_tokens(text)

    # ── Model Discovery ────────────────────────────────────────────

    async def list_models(self) -> list[str]:
        """Fetch available model names from Ollama's ``/api/tags`` endpoint."""
        if self._client is None or self._mock:
            return []

        try:
            response = await self._client.get(
                f"{self._base_url}/api/tags",
                timeout=self._timeout,
            )
            response.raise_for_status()
            data = response.json()
            return [model["name"] for model in data.get("models", [])]
        except Exception as exc:
            _logger.warning("Failed to list Ollama models: %s", exc)
            return []

    # ── Internal ───────────────────────────────────────────────────

    def _init_client(self) -> None:
        try:
            import httpx

            self._client = httpx.AsyncClient()
        except ImportError:
            _logger.warning("httpx package not installed — using mock fallback")
            self._client = None

    def _build_payload(self, request: CompletionRequest) -> dict[str, Any]:
        payload: dict[str, Any] = {
            "model": request.model,
            "messages": [{"role": m.role, "content": m.content} for m in request.messages],
        }
        options: dict[str, Any] = {}
        if request.max_tokens is not None:
            options["num_predict"] = request.max_tokens
        if request.temperature is not None:
            options["temperature"] = request.temperature
        if options:
            payload["options"] = options
        payload.update(request.extra)
        return payload
Methods:
list_models async
list_models() -> list[str]

Fetch available model names from Ollama's /api/tags endpoint.

Source code in src/forge/ai/adapters/ollama.py
async def list_models(self) -> list[str]:
    """Fetch available model names from Ollama's ``/api/tags`` endpoint."""
    if self._client is None or self._mock:
        return []

    try:
        response = await self._client.get(
            f"{self._base_url}/api/tags",
            timeout=self._timeout,
        )
        response.raise_for_status()
        data = response.json()
        return [model["name"] for model in data.get("models", [])]
    except Exception as exc:
        _logger.warning("Failed to list Ollama models: %s", exc)
        return []

OpenAIAdapter

Bases: BaseAdapter

Thin wrapper around the OpenAI Python SDK.

Requires openai to be installed. Falls back to mock behaviour when the library is unavailable.

Source code in src/forge/ai/adapters/openai.py
class OpenAIAdapter(BaseAdapter):
    """
    Thin wrapper around the OpenAI Python SDK.

    Requires ``openai`` to be installed.  Falls back to mock behaviour
    when the library is unavailable.
    """

    def __init__(self, api_key: str | None = None, timeout: float = 30.0) -> None:
        self._api_key = api_key
        self._timeout = timeout
        self._client: Any = None
        self._init_client()

    @property
    def provider(self) -> str:
        return "openai"

    # ── BaseAdapter ────────────────────────────────────────────────

    async def complete(self, request: CompletionRequest) -> CompletionResponse:
        if self._client is None or not self._api_key:
            return _mock_complete(request)

        payload = self._build_payload(request)
        try:
            response = await self._client.chat.completions.create(**payload)
        except Exception as exc:
            import openai

            if isinstance(exc, openai.RateLimitError):
                raise RateLimitError(f"OpenAI API rate limit exceeded: {exc}") from exc
            if isinstance(exc, openai.AuthenticationError):
                raise AIProviderError(f"OpenAI API authentication failed: {exc}") from exc
            if isinstance(exc, openai.NotFoundError):
                raise ModelNotFoundError(f"OpenAI model not found: {exc}") from exc
            raise AIProviderError(f"OpenAI API error: {exc}") from exc

        return CompletionResponse(
            model=response.model or request.model,
            message=Message(
                role="assistant",
                content=response.choices[0].message.content or "",
            ),
            usage=Usage(
                input_tokens=response.usage.prompt_tokens if response.usage else 0,
                output_tokens=response.usage.completion_tokens if response.usage else 0,
            ),
            provider="openai",
        )

    async def stream(self, request: CompletionRequest) -> AsyncIterator[StreamChunk]:
        if self._client is None or not self._api_key:
            async for chunk in _mock_stream(request):
                yield chunk
            return

        payload = self._build_payload(request)
        payload["stream"] = True
        payload["stream_options"] = {"include_usage": True}
        try:
            response = await self._client.chat.completions.create(**payload)
        except Exception as exc:
            import openai

            if isinstance(exc, openai.RateLimitError):
                raise RateLimitError(f"OpenAI API rate limit exceeded: {exc}") from exc
            if isinstance(exc, openai.AuthenticationError):
                raise AIProviderError(f"OpenAI API authentication failed: {exc}") from exc
            if isinstance(exc, openai.NotFoundError):
                raise ModelNotFoundError(f"OpenAI model not found: {exc}") from exc
            raise AIProviderError(f"OpenAI API error: {exc}") from exc

        try:
            async for event in response:
                delta_content = ""
                finish_reason = None
                usage = None

                if event.choices:
                    choice = event.choices[0]
                    if choice.delta and choice.delta.content:
                        delta_content = choice.delta.content
                    if choice.finish_reason:
                        finish_reason = choice.finish_reason

                if hasattr(event, "usage") and event.usage is not None:
                    usage = Usage(
                        input_tokens=event.usage.prompt_tokens,
                        output_tokens=event.usage.completion_tokens,
                    )

                if delta_content or finish_reason or usage:
                    yield StreamChunk(
                        delta=delta_content,
                        finish_reason=finish_reason,
                        usage=usage,
                        model=event.model or request.model,
                        provider="openai",
                    )
        except Exception as exc:
            raise StreamInterruptedError(f"OpenAI stream interrupted: {exc}") from exc

    def count_tokens(self, text: str) -> int:
        return TokenCounter.count_tokens(text)

    # ── Internal ───────────────────────────────────────────────────

    def _init_client(self) -> None:
        try:
            import openai

            if not self._api_key:
                self._client = None
            else:
                self._client = openai.AsyncOpenAI(api_key=self._api_key, timeout=self._timeout)
        except (ImportError, Exception) as exc:
            _logger.warning("Failed to initialize OpenAI client: %s — using mock", exc)
            self._client = None

    def _build_payload(self, request: CompletionRequest) -> dict[str, Any]:
        payload: dict[str, Any] = {
            "model": request.model,
            "messages": [{"role": m.role, "content": m.content} for m in request.messages],
        }
        if request.max_tokens is not None:
            payload["max_tokens"] = request.max_tokens
        if request.temperature is not None:
            payload["temperature"] = request.temperature
        payload.update(request.extra)
        return payload

RateLimitError

Bases: AIProviderError

Raised when the provider rate-limits the client.

Source code in src/forge/ai/exceptions.py
class RateLimitError(AIProviderError):
    """Raised when the provider rate-limits the client."""

StreamInterruptedError

Bases: AIError

Raised when an active stream is cut short or disconnected.

Source code in src/forge/ai/exceptions.py
class StreamInterruptedError(AIError):
    """Raised when an active stream is cut short or disconnected."""

StructuredOutputError

Bases: AIError

Raised when the model fails to produce output conforming to the Pydantic schema after the maximum retry attempts have been exhausted.

Source code in src/forge/ai/exceptions.py
class StructuredOutputError(AIError):
    """
    Raised when the model fails to produce output conforming to the Pydantic schema
    after the maximum retry attempts have been exhausted.
    """

    def __init__(
        self,
        schema_name: str,
        attempts: int,
        last_response: str,
        last_error: str,
    ) -> None:
        self.schema_name = schema_name
        self.attempts = attempts
        self.last_response = last_response
        self.last_error = last_error
        super().__init__(
            f"Failed to produce structured output conforming to '{schema_name}' "
            f"after {attempts} attempt(s). Last error: {last_error}. "
            f"Last raw response: {last_response!r}"
        )

TokenCounter

Estimates token usage and enforces pre-request budget limits.

Source code in src/forge/ai/tokens.py
class TokenCounter:
    """Estimates token usage and enforces pre-request budget limits."""

    @staticmethod
    def count_tokens(text: str) -> int:
        """Estimate the number of tokens in *text*."""
        return _chars_per_token(text)

    @staticmethod
    def count_messages(messages: list[Message]) -> int:
        """Estimate total tokens across a list of messages."""
        overhead = len(messages) * 4  # ~4 tokens per message for role markers
        content = sum(_chars_per_token(m.content) for m in messages)
        return overhead + content

    @staticmethod
    def estimate_cost(usage: Usage, model: str) -> float:
        """
        Calculate cost in USD from token counts using the pricing table.

        Returns 0.0 for unknown models.
        """
        prices = _PRICING.get(model)
        if prices is None:
            return 0.0
        input_cost = (usage.input_tokens / 1_000_000) * prices["input"]
        output_cost = (usage.output_tokens / 1_000_000) * prices["output"]
        return round(input_cost + output_cost, 6)

    @staticmethod
    def check_budget(
        request: CompletionRequest,
        max_tokens_limit: int,
    ) -> None:
        """
        Raise :class:`TokenLimitError` if the estimated request exceeds *max_tokens_limit*.

        Catches runaway prompts before calling the adapter.
        """
        estimated = TokenCounter.count_messages(request.messages)
        if estimated > max_tokens_limit:
            raise TokenLimitError(
                f"Estimated {estimated} input tokens exceeds limit of {max_tokens_limit}"
            )
Methods:
check_budget staticmethod
check_budget(request: CompletionRequest, max_tokens_limit: int) -> None

Raise :class:TokenLimitError if the estimated request exceeds max_tokens_limit.

Catches runaway prompts before calling the adapter.

Source code in src/forge/ai/tokens.py
@staticmethod
def check_budget(
    request: CompletionRequest,
    max_tokens_limit: int,
) -> None:
    """
    Raise :class:`TokenLimitError` if the estimated request exceeds *max_tokens_limit*.

    Catches runaway prompts before calling the adapter.
    """
    estimated = TokenCounter.count_messages(request.messages)
    if estimated > max_tokens_limit:
        raise TokenLimitError(
            f"Estimated {estimated} input tokens exceeds limit of {max_tokens_limit}"
        )
count_messages staticmethod
count_messages(messages: list[Message]) -> int

Estimate total tokens across a list of messages.

Source code in src/forge/ai/tokens.py
@staticmethod
def count_messages(messages: list[Message]) -> int:
    """Estimate total tokens across a list of messages."""
    overhead = len(messages) * 4  # ~4 tokens per message for role markers
    content = sum(_chars_per_token(m.content) for m in messages)
    return overhead + content
count_tokens staticmethod
count_tokens(text: str) -> int

Estimate the number of tokens in text.

Source code in src/forge/ai/tokens.py
@staticmethod
def count_tokens(text: str) -> int:
    """Estimate the number of tokens in *text*."""
    return _chars_per_token(text)
estimate_cost staticmethod
estimate_cost(usage: Usage, model: str) -> float

Calculate cost in USD from token counts using the pricing table.

Returns 0.0 for unknown models.

Source code in src/forge/ai/tokens.py
@staticmethod
def estimate_cost(usage: Usage, model: str) -> float:
    """
    Calculate cost in USD from token counts using the pricing table.

    Returns 0.0 for unknown models.
    """
    prices = _PRICING.get(model)
    if prices is None:
        return 0.0
    input_cost = (usage.input_tokens / 1_000_000) * prices["input"]
    output_cost = (usage.output_tokens / 1_000_000) * prices["output"]
    return round(input_cost + output_cost, 6)

TokenLimitError

Bases: AIError

Raised when the requested prompt exceeds the model's budget or maximum context.

Source code in src/forge/ai/exceptions.py
class TokenLimitError(AIError):
    """Raised when the requested prompt exceeds the model's budget or maximum context."""

Functions:

complete async

complete(messages: list[Message], model: str | None = None, max_tokens: int | None = None, temperature: float | None = None, output_schema: type[BaseModel] | None = None, fallback_models: list[str] | None = None, max_retries: int | None = None, **kwargs: Any) -> CompletionResponse | BaseModel

Convenience function to perform an AI completion request.

Uses the active ForgeRuntime context.

Source code in src/forge/ai/__init__.py
async def complete(
    messages: list[Message],
    model: str | None = None,
    max_tokens: int | None = None,
    temperature: float | None = None,
    output_schema: type[BaseModel] | None = None,
    fallback_models: list[str] | None = None,
    max_retries: int | None = None,
    **kwargs: Any,
) -> CompletionResponse | BaseModel:
    """
    Convenience function to perform an AI completion request.

    Uses the active ForgeRuntime context.
    """
    runtime = ForgeRuntime.get_active()
    ai_module: AIModule = runtime.get(AIModule)  # type: ignore[assignment]

    model_name = model or ai_module._config_ai.default_model
    request = CompletionRequest(
        model=model_name,
        messages=messages,
        max_tokens=max_tokens,
        temperature=temperature,
        extra=kwargs,
    )
    return await ai_module.complete(
        request=request,
        output_schema=output_schema,
        fallback_models=fallback_models,
        max_retries=max_retries,
    )

stream async

stream(messages: list[Message], model: str | None = None, max_tokens: int | None = None, temperature: float | None = None, **kwargs: Any) -> AsyncIterator[StreamChunk]

Convenience function to perform an AI streaming completion request.

Uses the active ForgeRuntime context.

Source code in src/forge/ai/__init__.py
async def stream(
    messages: list[Message],
    model: str | None = None,
    max_tokens: int | None = None,
    temperature: float | None = None,
    **kwargs: Any,
) -> AsyncIterator[StreamChunk]:
    """
    Convenience function to perform an AI streaming completion request.

    Uses the active ForgeRuntime context.
    """
    runtime = ForgeRuntime.get_active()
    ai_module: AIModule = runtime.get(AIModule)  # type: ignore[assignment]

    model_name = model or ai_module._config_ai.default_model
    request = CompletionRequest(
        model=model_name,
        messages=messages,
        max_tokens=max_tokens,
        temperature=temperature,
        extra=kwargs,
    )
    async for chunk in ai_module.stream(request):
        yield chunk