112 lines
3.4 KiB
Python
112 lines
3.4 KiB
Python
"""Runtime role and task-model resolution."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from dataclasses import dataclass
|
|
|
|
from scripts.lib.model_config import resolve_model_profile
|
|
|
|
|
|
ROLE_DEFAULTS = {
|
|
"dr_plan": {
|
|
"skills": ["document-ingest", "search-gateway", "search-strategy", "source-quality", "length-budget", "mckinsey-method"],
|
|
"temperature": 0.4,
|
|
"max_tokens": 12000,
|
|
"max_concurrency": 1,
|
|
},
|
|
"dr_pm": {
|
|
"skills": ["length-budget", "evidence-table", "mckinsey-method"],
|
|
"temperature": 0.2,
|
|
"max_tokens": 8000,
|
|
"max_concurrency": 1,
|
|
},
|
|
"dr_searcher": {
|
|
"skills": ["search-gateway", "search-strategy", "source-quality"],
|
|
"temperature": 0.1,
|
|
"max_tokens": 6000,
|
|
"max_concurrency": 6,
|
|
},
|
|
"dr_analyst": {
|
|
"skills": ["search-gateway", "search-strategy", "source-quality", "evidence-table", "mckinsey-method"],
|
|
"temperature": 0.3,
|
|
"max_tokens": 14000,
|
|
"max_concurrency": 6,
|
|
},
|
|
"dr_verifier": {
|
|
"skills": ["search-gateway", "search-strategy", "source-quality", "evidence-table"],
|
|
"temperature": 0.2,
|
|
"max_tokens": 10000,
|
|
"max_concurrency": 4,
|
|
},
|
|
"dr_chief_editor": {
|
|
"skills": ["mckinsey-method", "evidence-table", "output-hygiene"],
|
|
"temperature": 0.2,
|
|
"max_tokens": 16000,
|
|
"max_concurrency": 1,
|
|
},
|
|
"dr_editor_in_chief": {
|
|
"skills": ["mckinsey-method", "citation-manager", "humanizer-cn", "output-hygiene"],
|
|
"temperature": 0.4,
|
|
"max_tokens": 20000,
|
|
"max_concurrency": 1,
|
|
},
|
|
"dr_reporter": {
|
|
"skills": ["pdf-reportlab", "citation-manager", "output-hygiene"],
|
|
"temperature": 0.1,
|
|
"max_tokens": 6000,
|
|
"max_concurrency": 1,
|
|
},
|
|
}
|
|
|
|
|
|
@dataclass(frozen=True)
|
|
class RoleDefinition:
|
|
name: str
|
|
model: str
|
|
skills: list[str]
|
|
temperature: float
|
|
max_tokens: int
|
|
max_concurrency: int
|
|
|
|
|
|
class RuntimeProfile:
|
|
def __init__(self, *, profile: str, roles: dict[str, RoleDefinition], task_types: dict[str, str]) -> None:
|
|
self.profile = profile
|
|
self.roles = roles
|
|
self.task_types = task_types
|
|
|
|
def role_for_task(self, task_type: str) -> RoleDefinition:
|
|
role_name = self.task_types.get(task_type)
|
|
if not role_name:
|
|
raise KeyError(f"unknown task_type: {task_type}")
|
|
if role_name not in self.roles:
|
|
raise KeyError(f"task_type {task_type} maps to missing role {role_name}")
|
|
return self.roles[role_name]
|
|
|
|
|
|
def resolve_runtime_profile(
|
|
*,
|
|
profile: str | None = None,
|
|
overrides: dict[str, str] | None = None,
|
|
) -> RuntimeProfile:
|
|
resolved = resolve_model_profile(profile=profile, overrides=overrides)
|
|
role_models = resolved["roles"]
|
|
roles: dict[str, RoleDefinition] = {}
|
|
for name, defaults in ROLE_DEFAULTS.items():
|
|
model = role_models.get(name)
|
|
if not model:
|
|
continue
|
|
roles[name] = RoleDefinition(
|
|
name=name,
|
|
model=model,
|
|
skills=list(defaults["skills"]),
|
|
temperature=float(defaults["temperature"]),
|
|
max_tokens=int(defaults["max_tokens"]),
|
|
max_concurrency=int(defaults["max_concurrency"]),
|
|
)
|
|
return RuntimeProfile(
|
|
profile=resolved["profile"],
|
|
roles=roles,
|
|
task_types=dict(resolved.get("task_types") or {}),
|
|
)
|