Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
12 changes: 12 additions & 0 deletions src/powercontext/builtin/runtime/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,13 @@
InferenceConfig,
RuntimeConfig,
)
from powercontext.builtin.runtime.decision_model import (
DecisionModel,
DecisionModelOption,
DecisionModelRequest,
DecisionModelResult,
StructuredDecisionModel,
)
from powercontext.builtin.runtime.errors import InvalidRuntimeRequestError, TopicMemoryProcessingUnavailableError
from powercontext.builtin.runtime.models import (
ApproveArtifactCandidateRequest,
Expand Down Expand Up @@ -213,6 +220,10 @@
"ContextAssemblySection",
"CreateDreamRunRequest",
"DatabaseConfig",
"DecisionModel",
"DecisionModelOption",
"DecisionModelRequest",
"DecisionModelResult",
"DreamApplication",
"DreamRun",
"DreamRunPage",
Expand Down Expand Up @@ -337,6 +348,7 @@
"Statistics",
"StatisticsApplication",
"StatisticsPeriod",
"StructuredDecisionModel",
"SubmitSourceObservation",
"TopicMemoryApplication",
"TopicMemoryFlushResult",
Expand Down
39 changes: 35 additions & 4 deletions src/powercontext/builtin/runtime/config.py
Original file line number Diff line number Diff line change
Expand Up @@ -139,6 +139,7 @@ def reject_boolean_worker_quota(cls, value: Any) -> Any:
source_window_limit: int = Field(default=100, ge=1)
context_assembly_max_entries: int = Field(default=8, ge=1)
memory_extraction_profile: MemoryExtractionProfile = MemoryExtractionProfile.CODING
decision_assistance_enabled: bool = False
memory_rerank_enabled: bool = False
memory_rerank_candidate_limit: int = Field(default=30, ge=1, le=100)
recall_gate_enabled: bool = False
Expand Down Expand Up @@ -248,14 +249,20 @@ class InferenceConfig(BaseModel):
embedding_normalization: Literal["none", "unit"] = "unit"
embedding_timeout_seconds: float = Field(default=30.0, gt=0)
embedding_batch_size: int = Field(default=10, ge=1)
decision_model: str | None = None
decision_base_url: AnyHttpUrl | None = None
decision_headers: dict[str, SecretStr] = Field(default_factory=dict, repr=False)
decision_model_settings: dict[str, JsonValue] = Field(default_factory=dict)
decision_timeout_seconds: float | None = Field(default=None, gt=0)
decision_max_requests: int | None = Field(default=None, ge=1)
rerank_model: str | None = None
rerank_base_url: AnyHttpUrl | None = None
rerank_headers: dict[str, SecretStr] = Field(default_factory=dict, repr=False)
rerank_model_settings: dict[str, JsonValue] = Field(default_factory=dict)
rerank_timeout_seconds: float | None = Field(default=None, gt=0)
rerank_max_requests: int | None = Field(default=None, ge=1)

@field_validator("generation_model", "embedding_model", "embedding_profile_id", "rerank_model")
@field_validator("generation_model", "embedding_model", "embedding_profile_id", "decision_model", "rerank_model")
@classmethod
def validate_optional_identifier(cls, value: str | None) -> str | None:
if value is None:
Expand All @@ -275,7 +282,7 @@ def validate_normalization(cls, value: object) -> object:
raise ValueError("embedding normalization must be 'none' or 'unit'") # noqa: TRY003
return normalized

@field_validator("generation_headers", "embedding_headers", "rerank_headers")
@field_validator("generation_headers", "embedding_headers", "decision_headers", "rerank_headers")
@classmethod
def validate_headers(cls, value: dict[str, SecretStr]) -> dict[str, SecretStr]:
normalized_names: set[str] = set()
Expand All @@ -290,7 +297,12 @@ def validate_headers(cls, value: dict[str, SecretStr]) -> dict[str, SecretStr]:
normalized_names.add(normalized_name)
return value

@field_validator("generation_model_settings", "embedding_model_settings", "rerank_model_settings")
@field_validator(
"generation_model_settings",
"embedding_model_settings",
"decision_model_settings",
"rerank_model_settings",
)
@classmethod
def reserve_headers_field(cls, value: dict[str, JsonValue]) -> dict[str, JsonValue]:
if "extra_headers" in value:
Expand All @@ -310,14 +322,32 @@ def validate_embedding_profile(self) -> Self:

@model_validator(mode="after")
def validate_workload_overrides(self) -> Self:
self._validate_generation_overrides()
self._validate_embedding_overrides()
self._validate_decision_overrides()
self._validate_rerank_overrides()
return self

def _validate_generation_overrides(self) -> None:
if self.generation_model is None and self.generation_model_settings:
raise ValueError("generation_model_settings requires generation_model") # noqa: TRY003
if self.generation_model is None and (self.generation_base_url is not None or self.generation_headers):
raise ValueError("generation overrides require generation_model") # noqa: TRY003
self._validate_generation_budget()

def _validate_embedding_overrides(self) -> None:
if self.embedding_model is None and (
self.embedding_base_url is not None or self.embedding_headers or self.embedding_model_settings
):
raise ValueError("embedding overrides require a complete embedding profile") # noqa: TRY003

def _validate_decision_overrides(self) -> None:
if self.decision_base_url is not None and self.decision_model is None:
raise ValueError("decision_base_url requires decision_model") # noqa: TRY003
if self.decision_model is None and (self.decision_headers or self.decision_model_settings):
raise ValueError("decision overrides require decision_model") # noqa: TRY003

def _validate_rerank_overrides(self) -> None:
if self.rerank_base_url is not None and self.rerank_model is None:
raise ValueError("rerank_base_url requires rerank_model") # noqa: TRY003
if (
Expand All @@ -326,6 +356,8 @@ def validate_workload_overrides(self) -> Self:
and (self.rerank_headers or self.rerank_model_settings)
):
raise ValueError("rerank overrides require rerank_model or generation_model") # noqa: TRY003

def _validate_generation_budget(self) -> None:
max_tokens = self.generation_model_settings.get("max_tokens")
if max_tokens is not None and (
not isinstance(max_tokens, int) or isinstance(max_tokens, bool) or max_tokens < 1
Expand All @@ -343,7 +375,6 @@ def validate_workload_overrides(self) -> Self:
raise ValueError( # noqa: TRY003
f"Topic Memory generation budget is invalid: {error.error_code}"
) from error
return self


class ExternalSkillsConfig(BaseModel):
Expand Down
161 changes: 161 additions & 0 deletions src/powercontext/builtin/runtime/decision_model.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,161 @@
# Copyright (c) 2026 OceanBase.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.

"""Provider-neutral decision model contracts for runtime gates."""

from __future__ import annotations

from collections.abc import Mapping
from typing import Annotated, Protocol

from pydantic import BaseModel, ConfigDict, Field, JsonValue, field_validator, model_validator

from powercontext.builtin.inference import InferenceError, InferenceUsage, StructuredGenerator


class DecisionModelOption(BaseModel):
"""One selectable answer exposed to a deterministic runtime decision."""

model_config = ConfigDict(extra="forbid", frozen=True)

option_id: Annotated[str, Field(min_length=1, max_length=128)]
label: Annotated[str, Field(min_length=1, max_length=512)]
description: Annotated[str | None, Field(min_length=1, max_length=4096)] = None

@field_validator("option_id", "label", "description")
@classmethod
def validate_trimmed_text(cls, value: str | None) -> str | None:
if value is None:
return None
if value != value.strip():
raise ValueError("DecisionModel text fields must be trimmed") # noqa: TRY003
return value


class DecisionModelRequest(BaseModel):
"""A bounded, explicit-choice decision made outside the main agent loop."""

model_config = ConfigDict(extra="forbid", frozen=True)

operation: Annotated[str, Field(min_length=1, max_length=128)]
question: Annotated[str, Field(min_length=1, max_length=8192)]
options: Annotated[tuple[DecisionModelOption, ...], Field(min_length=2, max_length=128)]
context: Annotated[str | None, Field(min_length=1, max_length=32768)] = None
metadata: Mapping[str, JsonValue] = Field(default_factory=dict)

@field_validator("operation", "question", "context")
@classmethod
def validate_trimmed_text(cls, value: str | None) -> str | None:
if value is None:
return None
if value != value.strip():
raise ValueError("DecisionModel text fields must be trimmed") # noqa: TRY003
return value

@model_validator(mode="after")
def validate_option_ids(self) -> DecisionModelRequest:
option_ids = [option.option_id for option in self.options]
if len(option_ids) != len(set(option_ids)):
raise ValueError("DecisionModel option IDs must be unique") # noqa: TRY003
return self


class DecisionModelResult(BaseModel):
"""A selected option plus portable usage and fallback metadata."""

model_config = ConfigDict(extra="forbid", frozen=True)

selected_option_id: Annotated[str | None, Field(min_length=1, max_length=128)] = None
confidence: Annotated[float | None, Field(ge=0.0, le=1.0)] = None
scores: Mapping[str, Annotated[float, Field(ge=0.0, le=1.0)]] = Field(default_factory=dict)
rationale: Annotated[str | None, Field(min_length=1, max_length=4096)] = None
used_fallback: bool = False
fallback_reason: Annotated[str | None, Field(min_length=1, max_length=256)] = None
usage: InferenceUsage = Field(default_factory=lambda: InferenceUsage(requests=0))

@field_validator("selected_option_id", "rationale", "fallback_reason")
@classmethod
def validate_trimmed_text(cls, value: str | None) -> str | None:
if value is None:
return None
if value != value.strip():
raise ValueError("DecisionModel text fields must be trimmed") # noqa: TRY003
return value

@model_validator(mode="after")
def validate_fallback(self) -> DecisionModelResult:
if self.used_fallback:
if self.selected_option_id is not None:
raise ValueError("DecisionModel fallback results must not select an option") # noqa: TRY003
if self.fallback_reason is None:
raise ValueError("DecisionModel fallback results require fallback_reason") # noqa: TRY003
elif self.fallback_reason is not None:
raise ValueError("DecisionModel fallback_reason requires used_fallback") # noqa: TRY003
return self


class DecisionModel(Protocol):
"""Evaluate explicit-choice runtime decisions without exposing provider SDKs."""

policy_id: str

async def evaluate(self, request: DecisionModelRequest, /) -> DecisionModelResult:
"""Return one provider-neutral decision result or an explicit fallback."""

...


class StructuredDecisionModel:
"""Adapt a structured generator into a fail-open DecisionModel."""

def __init__(
self,
generator: StructuredGenerator[DecisionModelRequest, DecisionModelResult],
*,
policy_id: str,
) -> None:
self._generator = generator
self.policy_id = policy_id

async def evaluate(self, request: DecisionModelRequest, /) -> DecisionModelResult:
try:
result = await self._generator.generate(request)
except InferenceError as error:
return DecisionModelResult(used_fallback=True, fallback_reason=type(error).__name__)
output = result.output.model_copy(update={"usage": result.usage})
option_ids = {option.option_id for option in request.options}
if output.used_fallback:
return output
if output.selected_option_id not in option_ids:
return DecisionModelResult(
used_fallback=True,
fallback_reason="invalid_output",
usage=result.usage,
)
if any(option_id not in option_ids for option_id in output.scores):
return DecisionModelResult(
used_fallback=True,
fallback_reason="invalid_output",
usage=result.usage,
)
return output


__all__ = [
"DecisionModel",
"DecisionModelOption",
"DecisionModelRequest",
"DecisionModelResult",
"StructuredDecisionModel",
]
Loading
Loading