SkillOpt v0.1.0: initial release

- Skill optimization framework with training loop analogy
- 11 benchmarks, 4 model backends (Azure OpenAI, Claude, Codex, Qwen)
- WebUI for browser-based training control
- Pluggable architecture for extending benchmarks and backends
This commit is contained in:
CharlesYang030
2026-05-21 17:22:04 +00:00
commit 244e346b83
237 changed files with 30248 additions and 0 deletions
+403
View File
@@ -0,0 +1,403 @@
"""ReflACT model API with runtime backend selection for the student path."""
from __future__ import annotations
from typing import Any
from skillopt.model import azure_openai as _openai
from skillopt.model import claude_backend as _claude
from skillopt.model import qwen_backend as _qwen
from skillopt.model.backend_config import ( # noqa: F401
configure_claude_code_exec,
configure_codex_exec,
get_claude_code_exec_config,
get_codex_exec_config,
get_student_backend,
get_teacher_backend,
is_student_chat_backend,
is_student_exec_backend,
is_teacher_chat_backend,
set_student_backend,
set_teacher_backend,
)
def set_backend(name: str | None) -> str:
"""Backward-compatible global backend setter.
Historically the codebase used one shared backend for both teacher and
student. Keep that entry point so older scripts continue to work, while
mapping it onto the split teacher/student backend model.
"""
normalized = str(name or "azure_openai").strip().lower()
if normalized in {"azure_openai", "openai_chat", "azure", "azure-openai"}:
set_teacher_backend("openai_chat")
set_student_backend("openai_chat")
return "azure_openai"
if normalized in {"claude", "claude_chat", "anthropic"}:
set_teacher_backend("claude_chat")
set_student_backend("claude_chat")
return "claude_chat"
if normalized == "codex":
set_teacher_backend("openai_chat")
set_student_backend("codex_exec")
return "codex"
if normalized in {"codex_exec", "claude_code_exec"}:
set_teacher_backend("openai_chat")
set_student_backend(normalized)
return normalized
if normalized in {"qwen", "qwen_chat"}:
set_teacher_backend("openai_chat")
set_student_backend("qwen_chat")
return "qwen_chat"
raise ValueError(f"Unsupported legacy backend: {name!r}")
def get_backend_name() -> str:
"""Best-effort backward-compatible backend summary."""
teacher = get_teacher_backend()
student = get_student_backend()
if teacher == "claude_chat" and student == "claude_chat":
return "claude_chat"
if teacher == "openai_chat" and student == "openai_chat":
return "azure_openai"
if teacher == "openai_chat" and student == "codex_exec":
return "codex"
if teacher == "openai_chat" and student == "qwen_chat":
return "qwen_chat"
return f"{teacher}+{student}"
def chat_teacher(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
if get_teacher_backend() == "claude_chat":
return _claude.chat_teacher(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
return _openai.chat_teacher(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
timeout=timeout,
)
def chat_student(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
if get_student_backend() == "claude_chat":
return _claude.chat_student(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
if get_student_backend() == "qwen_chat":
return _qwen.chat_student(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
)
if not is_student_chat_backend():
raise NotImplementedError(
"chat_student is only supported with student_backend=openai_chat, claude_chat, or qwen_chat. "
"Exec backends are handled in environment-specific rollout code."
)
return _openai.chat_student(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
timeout=timeout,
)
def chat_teacher_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
if get_teacher_backend() == "claude_chat":
return _claude.chat_teacher_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
return _openai.chat_teacher_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_student_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
if get_student_backend() == "claude_chat":
return _claude.chat_student_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
if get_student_backend() == "qwen_chat":
return _qwen.chat_student_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
)
if not is_student_chat_backend():
raise NotImplementedError(
"chat_student_messages is only supported with student_backend=openai_chat, claude_chat, or qwen_chat. "
"Exec backends are handled in environment-specific rollout code."
)
return _openai.chat_student_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_messages_with_deployment(
deployment: str,
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
return _openai.chat_messages_with_deployment(
deployment=deployment,
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_with_deployment(
deployment: str,
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
return _openai.chat_with_deployment(
deployment=deployment,
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
reasoning_effort=reasoning_effort,
timeout=timeout,
)
def get_token_summary() -> dict:
summary = _openai.get_token_summary()
claude_summary = _claude.get_token_summary()
for stage, values in claude_summary.items():
if stage == "_total":
continue
if stage not in summary:
summary[stage] = values
continue
summary[stage]["calls"] += values["calls"]
summary[stage]["prompt_tokens"] += values["prompt_tokens"]
summary[stage]["completion_tokens"] += values["completion_tokens"]
summary[stage]["total_tokens"] += values["total_tokens"]
qwen_summary = _qwen.get_token_summary()
for stage, values in qwen_summary.items():
if stage == "_total":
continue
if stage not in summary:
summary[stage] = values
continue
summary[stage]["calls"] += values["calls"]
summary[stage]["prompt_tokens"] += values["prompt_tokens"]
summary[stage]["completion_tokens"] += values["completion_tokens"]
summary[stage]["total_tokens"] += values["total_tokens"]
total = {
"calls": 0,
"prompt_tokens": 0,
"completion_tokens": 0,
"total_tokens": 0,
}
for stage, values in summary.items():
if stage == "_total":
continue
total["calls"] += values["calls"]
total["prompt_tokens"] += values["prompt_tokens"]
total["completion_tokens"] += values["completion_tokens"]
total["total_tokens"] += values["total_tokens"]
summary["_total"] = total
return summary
def reset_token_tracker() -> None:
_openai.reset_token_tracker()
_claude.reset_token_tracker()
_qwen.reset_token_tracker()
def configure_azure_openai(
*,
endpoint: str | None = None,
api_version: str | None = None,
api_key: str | None = None,
auth_mode: str | None = None,
ad_scope: str | None = None,
managed_identity_client_id: str | None = None,
teacher_endpoint: str | None = None,
teacher_api_version: str | None = None,
teacher_api_key: str | None = None,
teacher_auth_mode: str | None = None,
teacher_ad_scope: str | None = None,
teacher_managed_identity_client_id: str | None = None,
student_endpoint: str | None = None,
student_api_version: str | None = None,
student_api_key: str | None = None,
student_auth_mode: str | None = None,
student_ad_scope: str | None = None,
student_managed_identity_client_id: str | None = None,
) -> None:
_openai.configure_azure_openai(
endpoint=endpoint,
api_version=api_version,
api_key=api_key,
auth_mode=auth_mode,
ad_scope=ad_scope,
managed_identity_client_id=managed_identity_client_id,
teacher_endpoint=teacher_endpoint,
teacher_api_version=teacher_api_version,
teacher_api_key=teacher_api_key,
teacher_auth_mode=teacher_auth_mode,
teacher_ad_scope=teacher_ad_scope,
teacher_managed_identity_client_id=teacher_managed_identity_client_id,
student_endpoint=student_endpoint,
student_api_version=student_api_version,
student_api_key=student_api_key,
student_auth_mode=student_auth_mode,
student_ad_scope=student_ad_scope,
student_managed_identity_client_id=student_managed_identity_client_id,
)
def configure_qwen_chat(
*,
base_url: str | None = None,
api_key: str | None = None,
temperature: float | str | None = None,
timeout_seconds: float | str | None = None,
max_tokens: int | str | None = None,
enable_thinking: bool | str | None = None,
) -> None:
_qwen.configure_qwen_chat(
base_url=base_url,
api_key=api_key,
temperature=temperature,
timeout_seconds=timeout_seconds,
max_tokens=max_tokens,
enable_thinking=enable_thinking,
)
def set_reasoning_effort(effort: str | None) -> None:
_openai.set_reasoning_effort(effort)
_claude.set_reasoning_effort(effort)
_qwen.set_reasoning_effort(effort)
def set_student_deployment(deployment: str) -> None:
_openai.set_student_deployment(deployment)
_claude.set_student_deployment(deployment)
_qwen.set_student_deployment(deployment)
def set_teacher_deployment(deployment: str) -> None:
_openai.set_teacher_deployment(deployment)
_claude.set_teacher_deployment(deployment)
+881
View File
@@ -0,0 +1,881 @@
"""ReflACT Model backend — Azure OpenAI wrapper with token tracking.
Provides teacher/student dual-deployment chat functions and a global
TokenTracker for per-stage cost accounting. Previously llm/azure_openai.py.
"""
from __future__ import annotations
import json
import os
import subprocess
import threading
import time
from types import SimpleNamespace
from typing import Any
from openai import AzureOpenAI, OpenAI
# ── Configuration ─────────────────────────────────────────────────────────────
ENDPOINT = os.environ.get(
"AZURE_OPENAI_ENDPOINT",
"", # Set via env var or config: e.g. "https://your-resource.openai.azure.com/"
)
API_VERSION = os.environ.get("AZURE_OPENAI_API_VERSION", "2024-12-01-preview")
API_KEY = os.environ.get(
"AZURE_OPENAI_API_KEY",
"",
)
AUTH_MODE = os.environ.get("AZURE_OPENAI_AUTH_MODE", "azure_cli").strip().lower()
AD_SCOPE = os.environ.get(
"AZURE_OPENAI_AD_SCOPE",
"https://cognitiveservices.azure.com/.default",
)
MANAGED_IDENTITY_CLIENT_ID = os.environ.get(
"AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID",
"",
).strip()
TEACHER_ENDPOINT = (
os.environ.get("TEACHER_AZURE_OPENAI_ENDPOINT")
or os.environ.get("AZURE_OPENAI_TEACHER_ENDPOINT")
or ENDPOINT
)
STUDENT_ENDPOINT = (
os.environ.get("STUDENT_AZURE_OPENAI_ENDPOINT")
or os.environ.get("AZURE_OPENAI_STUDENT_ENDPOINT")
or ENDPOINT
)
TEACHER_API_VERSION = (
os.environ.get("TEACHER_AZURE_OPENAI_API_VERSION")
or os.environ.get("AZURE_OPENAI_TEACHER_API_VERSION")
or API_VERSION
)
STUDENT_API_VERSION = (
os.environ.get("STUDENT_AZURE_OPENAI_API_VERSION")
or os.environ.get("AZURE_OPENAI_STUDENT_API_VERSION")
or API_VERSION
)
TEACHER_API_KEY = (
os.environ.get("TEACHER_AZURE_OPENAI_API_KEY")
or os.environ.get("AZURE_OPENAI_TEACHER_API_KEY")
or API_KEY
)
STUDENT_API_KEY = (
os.environ.get("STUDENT_AZURE_OPENAI_API_KEY")
or os.environ.get("AZURE_OPENAI_STUDENT_API_KEY")
or API_KEY
)
TEACHER_AUTH_MODE = (
os.environ.get("TEACHER_AZURE_OPENAI_AUTH_MODE")
or os.environ.get("AZURE_OPENAI_TEACHER_AUTH_MODE")
or AUTH_MODE
).strip().lower()
STUDENT_AUTH_MODE = (
os.environ.get("STUDENT_AZURE_OPENAI_AUTH_MODE")
or os.environ.get("AZURE_OPENAI_STUDENT_AUTH_MODE")
or AUTH_MODE
).strip().lower()
TEACHER_AD_SCOPE = (
os.environ.get("TEACHER_AZURE_OPENAI_AD_SCOPE")
or os.environ.get("AZURE_OPENAI_TEACHER_AD_SCOPE")
or AD_SCOPE
)
STUDENT_AD_SCOPE = (
os.environ.get("STUDENT_AZURE_OPENAI_AD_SCOPE")
or os.environ.get("AZURE_OPENAI_STUDENT_AD_SCOPE")
or AD_SCOPE
)
TEACHER_MANAGED_IDENTITY_CLIENT_ID = (
os.environ.get("TEACHER_AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID")
or os.environ.get("AZURE_OPENAI_TEACHER_MANAGED_IDENTITY_CLIENT_ID")
or MANAGED_IDENTITY_CLIENT_ID
).strip()
STUDENT_MANAGED_IDENTITY_CLIENT_ID = (
os.environ.get("STUDENT_AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID")
or os.environ.get("AZURE_OPENAI_STUDENT_MANAGED_IDENTITY_CLIENT_ID")
or MANAGED_IDENTITY_CLIENT_ID
).strip()
TEACHER_DEPLOYMENT = os.environ.get("TEACHER_DEPLOYMENT", "gpt-5.5")
STUDENT_DEPLOYMENT = os.environ.get("STUDENT_DEPLOYMENT", "gpt-5.5")
REASONING_EFFORT: str | None = None
_AZ_CLI_TOKEN_CACHE: dict[str, dict[str, Any]] = {}
# Deployments that require Responses API
_RESPONSES_API_MODELS = {
"gpt-5.3-codex", "gpt-5.1-codex", "gpt-5.2-codex",
"gpt-5-codex", "codex-mini", "gpt-5.4-pro",
}
# ── Token Tracker ─────────────────────────────────────────────────────────────
class TokenTracker:
"""Thread-safe per-stage token counter."""
def __init__(self) -> None:
self._lock = threading.Lock()
self._data: dict[str, dict] = {}
def record(
self, stage: str, prompt_tokens: int, completion_tokens: int,
) -> None:
with self._lock:
if stage not in self._data:
self._data[stage] = {
"calls": 0,
"prompt_tokens": 0,
"completion_tokens": 0,
}
d = self._data[stage]
d["calls"] += 1
d["prompt_tokens"] += prompt_tokens
d["completion_tokens"] += completion_tokens
def summary(self) -> dict:
with self._lock:
out: dict = {}
total_p = total_c = total_calls = 0
for stage, d in sorted(self._data.items()):
out[stage] = {
"calls": d["calls"],
"prompt_tokens": d["prompt_tokens"],
"completion_tokens": d["completion_tokens"],
"total_tokens": d["prompt_tokens"] + d["completion_tokens"],
}
total_p += d["prompt_tokens"]
total_c += d["completion_tokens"]
total_calls += d["calls"]
out["_total"] = {
"calls": total_calls,
"prompt_tokens": total_p,
"completion_tokens": total_c,
"total_tokens": total_p + total_c,
}
return out
def reset(self) -> None:
with self._lock:
self._data.clear()
def stage_snapshot(self, stage: str) -> dict:
"""Return a copy of one stage's counters (or zeros if not tracked yet)."""
with self._lock:
d = self._data.get(stage, {})
return {
"calls": d.get("calls", 0),
"prompt_tokens": d.get("prompt_tokens", 0),
"completion_tokens": d.get("completion_tokens", 0),
"total_tokens": d.get("prompt_tokens", 0) + d.get("completion_tokens", 0),
}
tracker = TokenTracker()
# ── Client management ─────────────────────────────────────────────────────────
_teacher_client: AzureOpenAI | None = None
_student_client: AzureOpenAI | None = None
_teacher_lock = threading.Lock()
_student_lock = threading.Lock()
def _role_config(role: str) -> dict[str, str]:
if role == "teacher":
return {
"endpoint": TEACHER_ENDPOINT,
"api_version": TEACHER_API_VERSION,
"api_key": TEACHER_API_KEY,
"auth_mode": TEACHER_AUTH_MODE,
"ad_scope": TEACHER_AD_SCOPE,
"managed_identity_client_id": TEACHER_MANAGED_IDENTITY_CLIENT_ID,
}
if role == "student":
return {
"endpoint": STUDENT_ENDPOINT,
"api_version": STUDENT_API_VERSION,
"api_key": STUDENT_API_KEY,
"auth_mode": STUDENT_AUTH_MODE,
"ad_scope": STUDENT_AD_SCOPE,
"managed_identity_client_id": STUDENT_MANAGED_IDENTITY_CLIENT_ID,
}
raise ValueError(f"Unknown Azure OpenAI client role: {role!r}")
def _make_token_provider(
auth_mode: str,
ad_scope: str,
managed_identity_client_id: str,
):
try:
from azure.identity import ( # type: ignore[import-not-found]
AzureCliCredential,
ManagedIdentityCredential,
get_bearer_token_provider,
)
except ImportError as e:
if auth_mode == "azure_cli":
return _make_azure_cli_token_provider(ad_scope)
raise ImportError(
"Azure AD auth requires azure-identity. Install it with `pip install azure-identity`."
) from e
if auth_mode in {"managed_identity", "aad", "azure_ad"}:
if managed_identity_client_id:
credential = ManagedIdentityCredential(client_id=managed_identity_client_id)
else:
credential = ManagedIdentityCredential()
elif auth_mode == "azure_cli":
credential = AzureCliCredential()
else:
raise ValueError(
"Unsupported Azure OpenAI auth mode "
f"{auth_mode!r}; expected api_key, managed_identity, azure_ad, aad, or azure_cli."
)
return get_bearer_token_provider(credential, ad_scope)
def _make_azure_cli_token_provider(ad_scope: str):
"""Return an Azure CLI token provider compatible with AzureOpenAI.
This fallback avoids requiring azure-identity in environments where `az`
is already logged in. The SDK calls this provider whenever it needs a
bearer token.
"""
resource = ad_scope.removesuffix("/.default")
def _provider() -> str:
now = int(time.time())
cache = _AZ_CLI_TOKEN_CACHE.setdefault(resource, {"token": "", "expires_on": 0})
cached = str(cache.get("token") or "")
expires_on = int(cache.get("expires_on") or 0)
if cached and expires_on - now > 300:
return cached
raw = subprocess.check_output(
[
"az",
"account",
"get-access-token",
"--resource",
resource,
"-o",
"json",
],
text=True,
stderr=subprocess.STDOUT,
)
payload = json.loads(raw)
token = str(payload["accessToken"])
cache["token"] = token
cache["expires_on"] = int(payload.get("expires_on") or now + 3000)
return token
return _provider
def _make_client(role: str) -> AzureOpenAI:
cfg = _role_config(role)
auth_mode = cfg["auth_mode"]
if auth_mode in {"api_key", "key"}:
if not cfg["api_key"]:
raise ValueError(
f"Azure OpenAI API key is not configured for {role}. "
"Set model.azure_openai_api_key in the config or export AZURE_OPENAI_API_KEY."
)
return AzureOpenAI(
api_version=cfg["api_version"],
azure_endpoint=cfg["endpoint"],
api_key=cfg["api_key"],
)
return AzureOpenAI(
api_version=cfg["api_version"],
azure_endpoint=cfg["endpoint"],
azure_ad_token_provider=_make_token_provider(
auth_mode,
cfg["ad_scope"],
cfg["managed_identity_client_id"],
),
)
def get_teacher_client() -> AzureOpenAI:
global _teacher_client
with _teacher_lock:
if _teacher_client is None:
_teacher_client = _make_client("teacher")
return _teacher_client
def get_student_client() -> AzureOpenAI | OpenAI:
global _student_client
with _student_lock:
if _student_client is None:
# When using qwen_chat backend, return an OpenAI client pointing to vLLM
from skillopt.model.backend_config import get_student_backend
if get_student_backend() == "qwen_chat":
from skillopt.model import qwen_backend as _qwen
_student_client = OpenAI(
base_url=_qwen.BASE_URL,
api_key=_qwen.API_KEY or "dummy",
)
else:
_student_client = _make_client("student")
return _student_client
def _needs_responses_api(deployment: str) -> bool:
dep = deployment.lower()
return any(dep == m or dep.startswith(m + "-") for m in _RESPONSES_API_MODELS)
# ── Core chat function ────────────────────────────────────────────────────────
def _chat_impl(
client: AzureOpenAI,
deployment: str,
system: str,
user: str,
max_completion_tokens: int,
retries: int,
stage: str,
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
"""Call LLM, track tokens, return (text, usage_dict)."""
last_err = None
usage_info = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
for attempt in range(retries):
try:
if _needs_responses_api(deployment):
kwargs: dict[str, Any] = {
"model": deployment,
"instructions": system,
"input": [{"role": "user", "content": user}],
"max_output_tokens": max_completion_tokens,
}
actual_effort = reasoning_effort or REASONING_EFFORT
if actual_effort:
kwargs["reasoning"] = {"effort": actual_effort}
if timeout is not None:
kwargs["timeout"] = timeout
resp = client.responses.create(**kwargs)
text = getattr(resp, "output_text", None) or ""
if not text:
for item in getattr(resp, "output", None) or []:
for part in getattr(item, "content", []):
if getattr(part, "type", "") == "output_text":
text = part.text or ""
if hasattr(resp, "usage") and resp.usage:
usage_info = {
"prompt_tokens": getattr(resp.usage, "input_tokens", 0) or 0,
"completion_tokens": getattr(resp.usage, "output_tokens", 0) or 0,
"total_tokens": (
(getattr(resp.usage, "input_tokens", 0) or 0)
+ (getattr(resp.usage, "output_tokens", 0) or 0)
),
}
else:
kwargs: dict[str, Any] = dict(
model=deployment,
messages=[
{"role": "system", "content": system},
{"role": "user", "content": user},
],
max_completion_tokens=max_completion_tokens,
)
actual_effort = reasoning_effort or REASONING_EFFORT
if actual_effort is not None:
kwargs["reasoning_effort"] = actual_effort
if timeout is not None:
kwargs["timeout"] = timeout
resp = client.chat.completions.create(**kwargs)
text = resp.choices[0].message.content or ""
if resp.usage:
usage_info = {
"prompt_tokens": resp.usage.prompt_tokens or 0,
"completion_tokens": resp.usage.completion_tokens or 0,
"total_tokens": resp.usage.total_tokens or 0,
}
tracker.record(
stage,
usage_info["prompt_tokens"],
usage_info["completion_tokens"],
)
return text, usage_info
except Exception as e: # noqa: BLE001
last_err = e
sleep = min(2 ** attempt, 30)
time.sleep(sleep)
raise RuntimeError(f"LLM call failed after {retries} retries: {last_err}")
def _chat_messages_impl(
client: AzureOpenAI,
deployment: str,
messages: list[dict[str, Any]],
max_completion_tokens: int,
retries: int,
stage: str,
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
"""Call the model with a pre-built message list."""
last_err = None
usage_info = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
for attempt in range(retries):
try:
if _needs_responses_api(deployment):
input_items, instructions = _messages_to_responses_input(messages)
kwargs: dict[str, Any] = {
"model": deployment,
"input": input_items,
"max_output_tokens": max_completion_tokens,
}
if instructions:
kwargs["instructions"] = instructions
actual_effort = reasoning_effort or REASONING_EFFORT
if actual_effort:
kwargs["reasoning"] = {"effort": actual_effort}
if tools:
kwargs["tools"] = [_chat_tool_to_responses_tool(tool) for tool in tools]
if tool_choice is not None:
kwargs["tool_choice"] = tool_choice
if timeout is not None:
kwargs["timeout"] = timeout
resp = client.responses.create(**kwargs)
message, text = _responses_to_chat_message(resp)
if hasattr(resp, "usage") and resp.usage:
usage_info = {
"prompt_tokens": getattr(resp.usage, "input_tokens", 0) or 0,
"completion_tokens": getattr(resp.usage, "output_tokens", 0) or 0,
"total_tokens": (
(getattr(resp.usage, "input_tokens", 0) or 0)
+ (getattr(resp.usage, "output_tokens", 0) or 0)
),
}
else:
kwargs = dict(
model=deployment,
messages=messages,
max_completion_tokens=max_completion_tokens,
)
actual_effort = reasoning_effort or REASONING_EFFORT
if tools:
kwargs["tools"] = tools
if tool_choice is not None:
kwargs["tool_choice"] = tool_choice
# Some models (e.g. gpt-5.5) don't support reasoning_effort with function tools
elif actual_effort is not None:
kwargs["reasoning_effort"] = actual_effort
if timeout is not None:
kwargs["timeout"] = timeout
resp = client.chat.completions.create(**kwargs)
message = resp.choices[0].message
text = message.content or ""
if resp.usage:
usage_info = {
"prompt_tokens": resp.usage.prompt_tokens or 0,
"completion_tokens": resp.usage.completion_tokens or 0,
"total_tokens": resp.usage.total_tokens or 0,
}
tracker.record(
stage,
usage_info["prompt_tokens"],
usage_info["completion_tokens"],
)
return (message if return_message else text), usage_info
except Exception as e: # noqa: BLE001
last_err = e
sleep = min(2 ** attempt, 30)
time.sleep(sleep)
raise RuntimeError(f"LLM message call failed after {retries} retries: {last_err}")
def _chat_tool_to_responses_tool(tool: dict[str, Any]) -> dict[str, Any]:
"""Convert a Chat Completions function tool to Responses API format."""
if tool.get("type") == "function" and isinstance(tool.get("function"), dict):
fn = tool["function"]
return {
"type": "function",
"name": fn.get("name", ""),
"description": fn.get("description", ""),
"parameters": fn.get("parameters", {"type": "object", "properties": {}}),
}
return tool
def _messages_to_responses_input(messages: list[dict[str, Any]]) -> tuple[list[dict[str, Any]], str]:
"""Convert chat-style messages, including tool results, to Responses input."""
instructions: list[str] = []
input_items: list[dict[str, Any]] = []
for message in messages:
role = message.get("role")
content = message.get("content") or ""
if role == "system":
if content:
instructions.append(str(content))
continue
if role == "tool":
input_items.append({
"type": "function_call_output",
"call_id": str(message.get("tool_call_id", "")),
"output": str(content),
})
continue
if role == "assistant":
if content:
input_items.append({"role": "assistant", "content": str(content)})
for tool_call in message.get("tool_calls") or []:
function = tool_call.get("function", {}) or {}
input_items.append({
"type": "function_call",
"call_id": str(tool_call.get("id", "")),
"name": str(function.get("name", "")),
"arguments": str(function.get("arguments", "{}") or "{}"),
})
continue
if role in {"user", "developer"}:
input_items.append({"role": "user", "content": str(content)})
return input_items, "\n\n".join(instructions)
def _responses_to_chat_message(resp: Any) -> tuple[Any, str]:
"""Convert Responses output into the subset of Chat message API we use."""
text = getattr(resp, "output_text", None) or ""
tool_calls: list[dict[str, Any]] = []
for item in getattr(resp, "output", None) or []:
item_type = getattr(item, "type", "")
if item_type == "function_call":
tool_calls.append({
"id": getattr(item, "call_id", "") or getattr(item, "id", ""),
"type": "function",
"function": {
"name": getattr(item, "name", ""),
"arguments": getattr(item, "arguments", "") or "{}",
},
})
elif item_type == "message" and not text:
content_parts = getattr(item, "content", []) or []
for part in content_parts:
if getattr(part, "type", "") == "output_text":
text += getattr(part, "text", "") or ""
return SimpleNamespace(content=text, tool_calls=tool_calls), text
# ── Public API ────────────────────────────────────────────────────────────────
def configure_azure_openai(
*,
endpoint: str | None = None,
api_version: str | None = None,
api_key: str | None = None,
auth_mode: str | None = None,
ad_scope: str | None = None,
managed_identity_client_id: str | None = None,
teacher_endpoint: str | None = None,
teacher_api_version: str | None = None,
teacher_api_key: str | None = None,
teacher_auth_mode: str | None = None,
teacher_ad_scope: str | None = None,
teacher_managed_identity_client_id: str | None = None,
student_endpoint: str | None = None,
student_api_version: str | None = None,
student_api_key: str | None = None,
student_auth_mode: str | None = None,
student_ad_scope: str | None = None,
student_managed_identity_client_id: str | None = None,
) -> None:
global ENDPOINT, API_VERSION, API_KEY, AUTH_MODE, AD_SCOPE, MANAGED_IDENTITY_CLIENT_ID
global TEACHER_ENDPOINT, TEACHER_API_VERSION, TEACHER_API_KEY, TEACHER_AUTH_MODE
global TEACHER_AD_SCOPE, TEACHER_MANAGED_IDENTITY_CLIENT_ID
global STUDENT_ENDPOINT, STUDENT_API_VERSION, STUDENT_API_KEY, STUDENT_AUTH_MODE
global STUDENT_AD_SCOPE, STUDENT_MANAGED_IDENTITY_CLIENT_ID
global _teacher_client, _student_client
def _clean(value: str | None, *, lower: bool = False) -> str | None:
if value is None:
return None
str_value = str(value).strip()
if not str_value:
return None
if lower:
str_value = str_value.lower()
return str_value
def _set(global_name: str, value: str | None, env_key: str) -> None:
if value is None:
return
globals()[global_name] = value
os.environ[env_key] = value
shared_endpoint = _clean(endpoint)
shared_api_version = _clean(api_version)
shared_api_key = _clean(api_key)
shared_auth_mode = _clean(auth_mode, lower=True)
shared_ad_scope = _clean(ad_scope)
shared_managed_identity_client_id = _clean(managed_identity_client_id)
_set("ENDPOINT", shared_endpoint, "AZURE_OPENAI_ENDPOINT")
_set("API_VERSION", shared_api_version, "AZURE_OPENAI_API_VERSION")
_set("API_KEY", shared_api_key, "AZURE_OPENAI_API_KEY")
_set("AUTH_MODE", shared_auth_mode, "AZURE_OPENAI_AUTH_MODE")
_set("AD_SCOPE", shared_ad_scope, "AZURE_OPENAI_AD_SCOPE")
_set(
"MANAGED_IDENTITY_CLIENT_ID",
shared_managed_identity_client_id,
"AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID",
)
resolved_teacher_endpoint = _clean(teacher_endpoint) or shared_endpoint
resolved_teacher_api_version = _clean(teacher_api_version) or shared_api_version
resolved_teacher_api_key = _clean(teacher_api_key) or shared_api_key
resolved_teacher_auth_mode = _clean(teacher_auth_mode, lower=True) or shared_auth_mode
resolved_teacher_ad_scope = _clean(teacher_ad_scope) or shared_ad_scope
resolved_teacher_mi = (
_clean(teacher_managed_identity_client_id)
or shared_managed_identity_client_id
)
resolved_student_endpoint = _clean(student_endpoint) or shared_endpoint
resolved_student_api_version = _clean(student_api_version) or shared_api_version
resolved_student_api_key = _clean(student_api_key) or shared_api_key
resolved_student_auth_mode = _clean(student_auth_mode, lower=True) or shared_auth_mode
resolved_student_ad_scope = _clean(student_ad_scope) or shared_ad_scope
resolved_student_mi = (
_clean(student_managed_identity_client_id)
or shared_managed_identity_client_id
)
_set("TEACHER_ENDPOINT", resolved_teacher_endpoint, "TEACHER_AZURE_OPENAI_ENDPOINT")
_set(
"TEACHER_API_VERSION",
resolved_teacher_api_version,
"TEACHER_AZURE_OPENAI_API_VERSION",
)
_set("TEACHER_API_KEY", resolved_teacher_api_key, "TEACHER_AZURE_OPENAI_API_KEY")
_set("TEACHER_AUTH_MODE", resolved_teacher_auth_mode, "TEACHER_AZURE_OPENAI_AUTH_MODE")
_set("TEACHER_AD_SCOPE", resolved_teacher_ad_scope, "TEACHER_AZURE_OPENAI_AD_SCOPE")
_set(
"TEACHER_MANAGED_IDENTITY_CLIENT_ID",
resolved_teacher_mi,
"TEACHER_AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID",
)
_set("STUDENT_ENDPOINT", resolved_student_endpoint, "STUDENT_AZURE_OPENAI_ENDPOINT")
_set(
"STUDENT_API_VERSION",
resolved_student_api_version,
"STUDENT_AZURE_OPENAI_API_VERSION",
)
_set("STUDENT_API_KEY", resolved_student_api_key, "STUDENT_AZURE_OPENAI_API_KEY")
_set("STUDENT_AUTH_MODE", resolved_student_auth_mode, "STUDENT_AZURE_OPENAI_AUTH_MODE")
_set("STUDENT_AD_SCOPE", resolved_student_ad_scope, "STUDENT_AZURE_OPENAI_AD_SCOPE")
_set(
"STUDENT_MANAGED_IDENTITY_CLIENT_ID",
resolved_student_mi,
"STUDENT_AZURE_OPENAI_MANAGED_IDENTITY_CLIENT_ID",
)
with _teacher_lock:
_teacher_client = None
with _student_lock:
_student_client = None
def chat_teacher(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
"""Call the teacher model. Returns (response_text, usage_dict)."""
return _chat_impl(
get_teacher_client(), TEACHER_DEPLOYMENT,
system, user, max_completion_tokens, retries, stage, reasoning_effort, timeout,
)
def chat_with_deployment(
deployment: str,
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
"""Call an arbitrary deployment using the shared Azure client."""
return _chat_impl(
get_teacher_client(),
deployment,
system,
user,
max_completion_tokens,
retries,
stage,
reasoning_effort,
timeout,
)
def chat_student(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
timeout: int | None = None,
) -> tuple[str, dict]:
"""Call the student model. Returns (response_text, usage_dict)."""
return _chat_impl(
get_student_client(), STUDENT_DEPLOYMENT,
system, user, max_completion_tokens, retries, stage, reasoning_effort, timeout,
)
def chat_teacher_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
"""Call the teacher model with a pre-built chat message list."""
return _chat_messages_impl(
get_teacher_client(),
TEACHER_DEPLOYMENT,
messages,
max_completion_tokens,
retries,
stage,
reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_messages_with_deployment(
deployment: str,
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
"""Call an arbitrary deployment with a pre-built chat message list."""
return _chat_messages_impl(
get_teacher_client(),
deployment,
messages,
max_completion_tokens,
retries,
stage,
reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_student_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict]:
"""Call the student model with a pre-built chat message list."""
return _chat_messages_impl(
get_student_client(),
STUDENT_DEPLOYMENT,
messages,
max_completion_tokens,
retries,
stage,
reasoning_effort,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def get_token_summary() -> dict:
"""Return per-stage and total token usage."""
return tracker.summary()
def reset_token_tracker() -> None:
tracker.reset()
def set_student_deployment(deployment: str) -> None:
"""Change student deployment at runtime."""
global _student_client, STUDENT_DEPLOYMENT
STUDENT_DEPLOYMENT = deployment
os.environ["STUDENT_DEPLOYMENT"] = deployment
os.environ["AZURE_OPENAI_DEPLOYMENT"] = deployment
with _student_lock:
_student_client = None
try:
import llm_client as _legacy
_legacy.DEPLOYMENT = deployment
_legacy._client = None
except Exception:
pass
def set_reasoning_effort(effort: str | None) -> None:
"""Set reasoning effort for all LLM calls. None = off."""
global REASONING_EFFORT
REASONING_EFFORT = effort if effort else None
def get_reasoning_effort() -> str | None:
"""Return the process-wide reasoning effort for direct Azure client users."""
return REASONING_EFFORT
def set_teacher_deployment(deployment: str) -> None:
"""Change teacher deployment at runtime."""
global _teacher_client, TEACHER_DEPLOYMENT
TEACHER_DEPLOYMENT = deployment
os.environ["TEACHER_DEPLOYMENT"] = deployment
with _teacher_lock:
_teacher_client = None
+185
View File
@@ -0,0 +1,185 @@
"""Runtime backend configuration for teacher/student model calls."""
from __future__ import annotations
import os
from skillopt.model.common import default_model_for_backend, normalize_backend_name
def _parse_bool(value: str | None, default: bool) -> bool:
if value is None:
return default
return str(value).strip().lower() in {"1", "true", "yes", "on"}
TEACHER_BACKEND = normalize_backend_name(os.environ.get("TEACHER_BACKEND", "openai_chat"))
STUDENT_BACKEND = normalize_backend_name(os.environ.get("STUDENT_BACKEND", "openai_chat"))
CODEX_EXEC_PATH = os.environ.get("CODEX_EXEC_PATH", "codex")
CODEX_EXEC_SANDBOX = os.environ.get("CODEX_EXEC_SANDBOX", "workspace-write")
CODEX_EXEC_PROFILE = os.environ.get("CODEX_EXEC_PROFILE", "")
CODEX_EXEC_FULL_AUTO = _parse_bool(os.environ.get("CODEX_EXEC_FULL_AUTO"), True)
CODEX_EXEC_REASONING_EFFORT = os.environ.get("CODEX_EXEC_REASONING_EFFORT", "none")
CODEX_EXEC_USE_SDK = os.environ.get("CODEX_EXEC_USE_SDK", "auto")
CODEX_EXEC_NETWORK_ACCESS = _parse_bool(os.environ.get("CODEX_EXEC_NETWORK_ACCESS"), False)
CODEX_EXEC_WEB_SEARCH = _parse_bool(os.environ.get("CODEX_EXEC_WEB_SEARCH"), False)
CODEX_EXEC_APPROVAL_POLICY = os.environ.get("CODEX_EXEC_APPROVAL_POLICY", "never")
CLAUDE_CODE_EXEC_PATH = os.environ.get("CLAUDE_CODE_EXEC_PATH", "claude")
CLAUDE_CODE_EXEC_PROFILE = os.environ.get("CLAUDE_CODE_EXEC_PROFILE", "")
CLAUDE_CODE_EXEC_USE_SDK = os.environ.get("CLAUDE_CODE_EXEC_USE_SDK", "auto")
CLAUDE_CODE_EXEC_EFFORT = os.environ.get("CLAUDE_CODE_EXEC_EFFORT", "medium")
def _parse_int(value: str | None, default: int) -> int:
if value is None:
return default
try:
return int(str(value).strip())
except ValueError:
return default
EXEC_EMPTY_RESPONSE_RETRIES = max(0, _parse_int(os.environ.get("EXEC_EMPTY_RESPONSE_RETRIES"), 1))
CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS = max(
0,
_parse_int(os.environ.get("CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS"), 16384),
)
def set_teacher_backend(backend: str) -> None:
global TEACHER_BACKEND
TEACHER_BACKEND = normalize_backend_name(backend or "openai_chat")
if TEACHER_BACKEND not in {"openai_chat", "claude_chat"}:
raise ValueError(
f"Unsupported teacher backend: {TEACHER_BACKEND!r}. "
"Supported values are 'openai_chat' and 'claude_chat'."
)
os.environ["TEACHER_BACKEND"] = TEACHER_BACKEND
def get_teacher_backend() -> str:
return TEACHER_BACKEND
def set_student_backend(backend: str) -> None:
global STUDENT_BACKEND
STUDENT_BACKEND = normalize_backend_name(backend or "openai_chat")
if STUDENT_BACKEND not in {"openai_chat", "claude_chat", "qwen_chat", "codex_exec", "claude_code_exec"}:
raise ValueError(
f"Unsupported student backend: {STUDENT_BACKEND!r}. "
"Supported values are 'openai_chat', 'claude_chat', 'qwen_chat', 'codex_exec', and 'claude_code_exec'."
)
os.environ["STUDENT_BACKEND"] = STUDENT_BACKEND
def get_student_backend() -> str:
return STUDENT_BACKEND
def is_student_exec_backend() -> bool:
return STUDENT_BACKEND in {"codex_exec", "claude_code_exec"}
def is_teacher_chat_backend() -> bool:
return TEACHER_BACKEND in {"openai_chat", "claude_chat"}
def is_student_chat_backend() -> bool:
return STUDENT_BACKEND in {"openai_chat", "claude_chat", "qwen_chat"}
def configure_codex_exec(
*,
path: str | None = None,
sandbox: str | None = None,
profile: str | None = None,
full_auto: bool | None = None,
reasoning_effort: str | None = None,
use_sdk: str | None = None,
network_access: bool | None = None,
web_search: bool | None = None,
approval_policy: str | None = None,
) -> None:
global CODEX_EXEC_PATH, CODEX_EXEC_SANDBOX, CODEX_EXEC_PROFILE, CODEX_EXEC_FULL_AUTO, CODEX_EXEC_REASONING_EFFORT, CODEX_EXEC_USE_SDK, CODEX_EXEC_NETWORK_ACCESS, CODEX_EXEC_WEB_SEARCH, CODEX_EXEC_APPROVAL_POLICY
if path is not None:
CODEX_EXEC_PATH = str(path).strip() or "codex"
os.environ["CODEX_EXEC_PATH"] = CODEX_EXEC_PATH
if sandbox is not None:
CODEX_EXEC_SANDBOX = str(sandbox).strip() or "workspace-write"
os.environ["CODEX_EXEC_SANDBOX"] = CODEX_EXEC_SANDBOX
if profile is not None:
CODEX_EXEC_PROFILE = str(profile).strip()
os.environ["CODEX_EXEC_PROFILE"] = CODEX_EXEC_PROFILE
if full_auto is not None:
CODEX_EXEC_FULL_AUTO = bool(full_auto)
os.environ["CODEX_EXEC_FULL_AUTO"] = "true" if CODEX_EXEC_FULL_AUTO else "false"
if reasoning_effort is not None:
CODEX_EXEC_REASONING_EFFORT = str(reasoning_effort).strip() or "none"
os.environ["CODEX_EXEC_REASONING_EFFORT"] = CODEX_EXEC_REASONING_EFFORT
if use_sdk is not None:
CODEX_EXEC_USE_SDK = str(use_sdk).strip().lower() or "auto"
os.environ["CODEX_EXEC_USE_SDK"] = CODEX_EXEC_USE_SDK
if network_access is not None:
CODEX_EXEC_NETWORK_ACCESS = bool(network_access)
os.environ["CODEX_EXEC_NETWORK_ACCESS"] = "true" if CODEX_EXEC_NETWORK_ACCESS else "false"
if web_search is not None:
CODEX_EXEC_WEB_SEARCH = bool(web_search)
os.environ["CODEX_EXEC_WEB_SEARCH"] = "true" if CODEX_EXEC_WEB_SEARCH else "false"
if approval_policy is not None:
CODEX_EXEC_APPROVAL_POLICY = str(approval_policy).strip() or "never"
os.environ["CODEX_EXEC_APPROVAL_POLICY"] = CODEX_EXEC_APPROVAL_POLICY
def get_codex_exec_config() -> dict[str, str | bool | int]:
return {
"path": CODEX_EXEC_PATH,
"sandbox": CODEX_EXEC_SANDBOX,
"profile": CODEX_EXEC_PROFILE,
"full_auto": CODEX_EXEC_FULL_AUTO,
"reasoning_effort": CODEX_EXEC_REASONING_EFFORT,
"use_sdk": CODEX_EXEC_USE_SDK,
"network_access": CODEX_EXEC_NETWORK_ACCESS,
"web_search": CODEX_EXEC_WEB_SEARCH,
"approval_policy": CODEX_EXEC_APPROVAL_POLICY,
"empty_response_retries": EXEC_EMPTY_RESPONSE_RETRIES,
}
def configure_claude_code_exec(
*,
path: str | None = None,
profile: str | None = None,
use_sdk: str | None = None,
effort: str | None = None,
max_thinking_tokens: int | str | None = None,
) -> None:
global CLAUDE_CODE_EXEC_PATH, CLAUDE_CODE_EXEC_PROFILE, CLAUDE_CODE_EXEC_USE_SDK, CLAUDE_CODE_EXEC_EFFORT, CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS
if path is not None:
CLAUDE_CODE_EXEC_PATH = str(path).strip() or "claude"
os.environ["CLAUDE_CODE_EXEC_PATH"] = CLAUDE_CODE_EXEC_PATH
if profile is not None:
CLAUDE_CODE_EXEC_PROFILE = str(profile).strip()
os.environ["CLAUDE_CODE_EXEC_PROFILE"] = CLAUDE_CODE_EXEC_PROFILE
if use_sdk is not None:
CLAUDE_CODE_EXEC_USE_SDK = str(use_sdk).strip().lower() or "auto"
os.environ["CLAUDE_CODE_EXEC_USE_SDK"] = CLAUDE_CODE_EXEC_USE_SDK
if effort is not None:
CLAUDE_CODE_EXEC_EFFORT = str(effort).strip().lower() or "medium"
os.environ["CLAUDE_CODE_EXEC_EFFORT"] = CLAUDE_CODE_EXEC_EFFORT
if max_thinking_tokens is not None:
CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS = max(
0,
_parse_int(str(max_thinking_tokens), 16384),
)
os.environ["CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS"] = str(CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS)
def get_claude_code_exec_config() -> dict[str, str | int]:
return {
"path": CLAUDE_CODE_EXEC_PATH,
"profile": CLAUDE_CODE_EXEC_PROFILE,
"use_sdk": CLAUDE_CODE_EXEC_USE_SDK,
"effort": CLAUDE_CODE_EXEC_EFFORT,
"max_thinking_tokens": CLAUDE_CODE_EXEC_MAX_THINKING_TOKENS,
"empty_response_retries": EXEC_EMPTY_RESPONSE_RETRIES,
}
+359
View File
@@ -0,0 +1,359 @@
"""Claude CLI chat backend for ReflACT."""
from __future__ import annotations
import base64
import json
import mimetypes
import os
import shutil
import subprocess
import tempfile
import time
from typing import Any
from urllib.parse import unquote, urlparse
from skillopt.model.common import CompatAssistantMessage, CompatToolCall, CompatToolFunction, default_model_for_backend, tracker
CLAUDE_BIN = os.environ.get("CLAUDE_CLI_BIN", "claude")
CLAUDE_PERMISSION_MODE = os.environ.get("CLAUDE_PERMISSION_MODE", "dontAsk")
CLAUDE_SETTING_SOURCES = os.environ.get("CLAUDE_SETTING_SOURCES", "user,project")
CLAUDE_ALLOW_ATTACHMENT_READ = os.environ.get("CLAUDE_ALLOW_ATTACHMENT_READ", "1").strip().lower() not in {"0", "false", "no"}
TEACHER_DEPLOYMENT = os.environ.get("TEACHER_DEPLOYMENT", "claude-sonnet-4-6")
STUDENT_DEPLOYMENT = os.environ.get("STUDENT_DEPLOYMENT", "claude-sonnet-4-6")
REASONING_EFFORT: str | None = None
_VALID_EFFORTS = {"low", "medium", "high", "xhigh", "max"}
def _parse_data_uri(url: str) -> tuple[bytes, str]:
header, data = url.split(",", 1)
mime = header[5:].split(";", 1)[0] or "image/png"
return base64.b64decode(data), mime
def _content_to_text(content: Any, attachments: list[dict[str, Any]], *, image_counter: int) -> tuple[str, int]:
if isinstance(content, str):
return content, image_counter
if not isinstance(content, list):
return str(content), image_counter
parts: list[str] = []
for item in content:
if not isinstance(item, dict):
continue
item_type = item.get("type")
if item_type == "text":
parts.append(str(item.get("text", "")))
continue
if item_type != "image_url":
continue
image_counter += 1
label = f"[Attached image {image_counter}]"
parts.append(label)
image_url = item.get("image_url", {}) or {}
url = str(image_url.get("url", "") or "")
if not url:
continue
if url.startswith("data:") and ";base64," in url:
data, mime = _parse_data_uri(url)
attachments.append({"bytes": data, "mime": mime, "label": label})
continue
if url.startswith("file://"):
parsed = urlparse(url)
path = unquote(parsed.path)
if path:
attachments.append({"path": path, "label": label})
continue
if os.path.exists(url):
attachments.append({"path": url, "label": label})
return "".join(parts), image_counter
def _simplify_tool_schemas(tools: list[dict[str, Any]] | None) -> list[dict[str, Any]]:
simplified: list[dict[str, Any]] = []
for tool in tools or []:
function = tool.get("function", tool)
simplified.append({
"name": function.get("name", ""),
"description": function.get("description", ""),
"parameters": function.get("parameters", {}),
})
return simplified
def _build_prompt_from_messages(messages: list[dict[str, Any]], *, tools: list[dict[str, Any]] | None = None, tool_choice: str | dict[str, Any] | None = None, structured_output: bool = False) -> tuple[str, str, list[dict[str, Any]]]:
system_parts: list[str] = []
history_parts: list[str] = []
attachments: list[dict[str, Any]] = []
image_counter = 0
def _history_line(label: str, body: str) -> str:
stripped = body.strip()
if not stripped:
return f"- {label}:"
indented = stripped.replace("\n", "\n ")
return f"- {label}: {indented}"
for message in messages:
role = str(message.get("role", "user"))
text, image_counter = _content_to_text(message.get("content", ""), attachments, image_counter=image_counter)
if role == "system":
if text.strip():
system_parts.append(text.strip())
continue
if role == "assistant":
block = _history_line("Assistant", text)
tool_calls = message.get("tool_calls") or []
if tool_calls:
simplified_calls = []
for tool_call in tool_calls:
function = tool_call.get("function", {}) or {}
simplified_calls.append({
"name": function.get("name", ""),
"arguments": function.get("arguments", "{}"),
})
block += "\n Compatibility tool requests:\n" + json.dumps(simplified_calls, ensure_ascii=False, indent=2)
history_parts.append(block)
continue
if role == "tool":
tool_call_id = str(message.get("tool_call_id", "") or "")
history_parts.append(_history_line(f"Tool result (tool_call_id={tool_call_id})", text))
continue
history_parts.append(_history_line(role.capitalize(), text))
prompt_parts: list[str] = []
if tools:
simplified_tools = _simplify_tool_schemas(tools)
prompt_parts.append("Available compatibility tools:\n" + json.dumps(simplified_tools, ensure_ascii=False, indent=2))
prompt_parts.append("Do not execute these compatibility tools yourself. If you need one, request it in `tool_calls`. Each `arguments` field must be a JSON string.")
if tool_choice == "required":
prompt_parts.append("Tool choice policy: you must request at least one compatibility tool.")
elif isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
function = tool_choice.get("function", {}) or {}
prompt_parts.append(f"Tool choice policy: you must request the compatibility tool `{function.get('name', '')}`.")
history_text = "\n".join(part for part in history_parts if part).strip()
if history_text:
prompt_parts.append("History:\n" + history_text)
if structured_output:
prompt_parts.append("Return only JSON matching the provided schema.")
if tools:
prompt_parts.append("Set `content` to the assistant-visible reply. Set `tool_calls` to an empty array when no compatibility tool is needed.")
else:
prompt_parts.append("Answer the latest user request.")
return "\n\n".join(part for part in system_parts if part).strip(), "\n\n".join(prompt_parts), attachments
def _copy_attachments_to_temp(attachments: list[dict[str, Any]], temp_dir: str) -> list[dict[str, str]]:
copied: list[dict[str, str]] = []
for index, attachment in enumerate(attachments, 1):
source_path = attachment.get("path")
if source_path:
source_path = str(source_path)
source_suffix = os.path.splitext(source_path)[1]
target_path = os.path.join(temp_dir, f"image_{index}{source_suffix or '.bin'}")
shutil.copyfile(source_path, target_path)
copied.append({"path": target_path, "label": str(attachment.get("label", ""))})
continue
mime = str(attachment.get("mime", "image/png"))
suffix = mimetypes.guess_extension(mime) or ".png"
target_path = os.path.join(temp_dir, f"image_{index}{suffix}")
with open(target_path, "wb") as f:
f.write(attachment.get("bytes", b"") or b"")
copied.append({"path": target_path, "label": str(attachment.get("label", ""))})
return copied
def _append_attachment_instructions(prompt: str, copied_attachments: list[dict[str, str]]) -> str:
if not copied_attachments or not CLAUDE_ALLOW_ATTACHMENT_READ:
return prompt
lines = [
"Attached image files:",
*[f"- {item['label'] or f'Attached image {index}'}: {item['path']}" for index, item in enumerate(copied_attachments, 1)],
"If you need to inspect an attached image, you may use the built-in `Read` tool on those listed paths only. Do not use built-in tools for any other purpose.",
]
return prompt.rstrip() + "\n\n" + "\n".join(lines)
def _usage_from_result(result_event: dict[str, Any] | None) -> dict[str, int]:
usage = (result_event or {}).get("usage", {}) or {}
input_tokens = int(usage.get("input_tokens", 0) or 0)
output_tokens = int(usage.get("output_tokens", 0) or 0)
return {
"prompt_tokens": input_tokens,
"completion_tokens": output_tokens,
"total_tokens": input_tokens + output_tokens,
}
def _extract_result(event_stream: list[dict[str, Any]]) -> tuple[str, dict[str, Any] | None]:
result_event = None
for event in reversed(event_stream):
if event.get("type") == "result":
result_event = event
break
if result_event is None:
raise RuntimeError("Claude backend did not return a result event.")
content = result_event.get("result") or result_event.get("content") or ""
return str(content), result_event
def _check_claude_error(stderr_text: str, model: str) -> None:
lowered = stderr_text.lower()
if "invalid api key" in lowered or "authentication" in lowered or "login" in lowered:
raise RuntimeError("Claude CLI is not logged in. Run `claude auth login` (or start `claude` and use `/login`) first.")
if "unknown model" in lowered or "not available" in lowered or "invalid model" in lowered:
default_model = default_model_for_backend("claude")
raise RuntimeError(f"Claude backend tried to use model {model!r}, but your current Claude CLI/account rejected it. Try an available Claude model such as {default_model!r}.")
def _normalize_reasoning_effort(effort: str | None) -> str | None:
normalized = str(effort or "").strip().lower()
if not normalized or normalized == "off":
return None
if normalized in _VALID_EFFORTS:
return normalized
return None
def _assistant_message_schema() -> dict[str, Any]:
return {
"type": "object",
"properties": {
"content": {"type": "string"},
"tool_calls": {
"type": "array",
"items": {
"type": "object",
"properties": {
"name": {"type": "string"},
"arguments": {"type": "string"},
},
"required": ["name", "arguments"],
"additionalProperties": False,
},
},
},
"required": ["content", "tool_calls"],
"additionalProperties": False,
}
def _assistant_message_schema_wrapper() -> str:
return json.dumps(_assistant_message_schema(), ensure_ascii=False)
def _run_claude_print(*, system: str, prompt: str, model: str, tools: list[dict[str, Any]] | None, tool_choice: str | dict[str, Any] | None, return_message: bool, timeout: int | None, attachments: list[dict[str, Any]] | None = None) -> tuple[str, dict[str, Any], dict[str, int]]:
effort = _normalize_reasoning_effort(REASONING_EFFORT)
with tempfile.TemporaryDirectory(prefix="skillopt_claude_") as temp_dir:
copied_attachments = _copy_attachments_to_temp(attachments or [], temp_dir)
prompt_for_cli = _append_attachment_instructions(prompt, copied_attachments)
cmd = [CLAUDE_BIN, "-p", "--output-format", "json", "--permission-mode", CLAUDE_PERMISSION_MODE, "--add-dir", temp_dir]
if model:
cmd.extend(["--model", model])
if CLAUDE_SETTING_SOURCES:
cmd.extend(["--setting-sources", CLAUDE_SETTING_SOURCES])
if system:
cmd.extend(["--append-system-prompt", system])
if effort:
cmd.extend(["--thinking", effort])
structured_output = bool(return_message)
if structured_output:
cmd.extend(["--schema", _assistant_message_schema_wrapper()])
proc = subprocess.run(cmd + [prompt_for_cli], capture_output=True, text=True, timeout=timeout or 300, cwd=temp_dir)
stderr_text = (proc.stderr or "").strip()
if proc.returncode != 0:
_check_claude_error(stderr_text, model)
raise RuntimeError(stderr_text or f"Claude CLI exited with code {proc.returncode}")
stream = []
for raw_line in (proc.stdout or "").splitlines():
raw_line = raw_line.strip()
if not raw_line:
continue
try:
stream.append(json.loads(raw_line))
except json.JSONDecodeError:
continue
raw_text, result_event = _extract_result(stream)
usage_info = _usage_from_result(result_event)
return raw_text, result_event or {}, usage_info
def _compat_message_from_payload(payload: Any) -> CompatAssistantMessage:
if not isinstance(payload, dict):
return CompatAssistantMessage(content=str(payload or ""), tool_calls=[])
content = str(payload.get("content", "") or "")
tool_calls: list[CompatToolCall] = []
for index, tool_call in enumerate(payload.get("tool_calls", []) or [], start=1):
name = str(tool_call.get("name", "") or "")
arguments = str(tool_call.get("arguments", "{}") or "{}")
tool_calls.append(CompatToolCall(id=f"claude_tool_{index}", function=CompatToolFunction(name=name, arguments=arguments)))
return CompatAssistantMessage(content=content, tool_calls=tool_calls)
def _call_messages(messages: list[dict[str, Any]], max_completion_tokens: int, retries: int, stage: str, *, tools: list[dict[str, Any]] | None = None, tool_choice: str | dict[str, Any] | None = None, return_message: bool = False, deployment: str | None = None, timeout: int | None = None) -> tuple[Any, dict[str, int]]:
del max_completion_tokens
system, prompt, attachments = _build_prompt_from_messages(messages, tools=tools, tool_choice=tool_choice, structured_output=return_message)
model = deployment or STUDENT_DEPLOYMENT
last_err = None
for attempt in range(retries):
try:
raw_text, payload, usage_info = _run_claude_print(system=system, prompt=prompt, model=model, tools=tools, tool_choice=tool_choice, return_message=return_message, timeout=timeout, attachments=attachments)
tracker.record(stage, usage_info["prompt_tokens"], usage_info["completion_tokens"])
if return_message:
return _compat_message_from_payload(payload.get("result", payload)), usage_info
return raw_text, usage_info
except Exception as e: # noqa: BLE001
last_err = e
time.sleep(min(2 ** attempt, 15))
raise RuntimeError(f"Claude backend failed after {retries} retries: {last_err}")
def chat_teacher(system: str, user: str, max_completion_tokens: int = 16384, retries: int = 5, stage: str = "teacher", timeout: int | None = None) -> tuple[str, dict[str, int]]:
messages = [{"role": "system", "content": system}, {"role": "user", "content": user}]
return _call_messages(messages, max_completion_tokens, retries, stage, deployment=TEACHER_DEPLOYMENT, timeout=timeout)
def chat_student(system: str, user: str, max_completion_tokens: int = 16384, retries: int = 5, stage: str = "student", timeout: int | None = None) -> tuple[str, dict[str, int]]:
messages = [{"role": "system", "content": system}, {"role": "user", "content": user}]
return _call_messages(messages, max_completion_tokens, retries, stage, deployment=STUDENT_DEPLOYMENT, timeout=timeout)
def chat_with_deployment(deployment: str, system: str, user: str, max_completion_tokens: int = 16384, retries: int = 5, stage: str = "custom", timeout: int | None = None) -> tuple[str, dict[str, int]]:
messages = [{"role": "system", "content": system}, {"role": "user", "content": user}]
return _call_messages(messages, max_completion_tokens, retries, stage, deployment=deployment, timeout=timeout)
def chat_teacher_messages(messages: list[dict[str, Any]], max_completion_tokens: int = 16384, retries: int = 5, stage: str = "teacher", *, tools: list[dict[str, Any]] | None = None, tool_choice: str | dict[str, Any] | None = None, return_message: bool = False, timeout: int | None = None) -> tuple[Any, dict[str, int]]:
return _call_messages(messages, max_completion_tokens, retries, stage, tools=tools, tool_choice=tool_choice, return_message=return_message, deployment=TEACHER_DEPLOYMENT, timeout=timeout)
def chat_student_messages(messages: list[dict[str, Any]], max_completion_tokens: int = 16384, retries: int = 5, stage: str = "student", *, tools: list[dict[str, Any]] | None = None, tool_choice: str | dict[str, Any] | None = None, return_message: bool = False, timeout: int | None = None) -> tuple[Any, dict[str, int]]:
return _call_messages(messages, max_completion_tokens, retries, stage, tools=tools, tool_choice=tool_choice, return_message=return_message, deployment=STUDENT_DEPLOYMENT, timeout=timeout)
def chat_messages_with_deployment(deployment: str, messages: list[dict[str, Any]], max_completion_tokens: int = 16384, retries: int = 5, stage: str = "custom", *, tools: list[dict[str, Any]] | None = None, tool_choice: str | dict[str, Any] | None = None, return_message: bool = False, timeout: int | None = None) -> tuple[Any, dict[str, int]]:
return _call_messages(messages, max_completion_tokens, retries, stage, tools=tools, tool_choice=tool_choice, return_message=return_message, deployment=deployment, timeout=timeout)
def get_token_summary() -> dict[str, dict[str, int]]:
return tracker.summary()
def reset_token_tracker() -> None:
tracker.reset()
def set_reasoning_effort(effort: str | None) -> None:
global REASONING_EFFORT
REASONING_EFFORT = effort if effort else None
def set_student_deployment(deployment: str) -> None:
global STUDENT_DEPLOYMENT
STUDENT_DEPLOYMENT = deployment or default_model_for_backend("claude")
os.environ["STUDENT_DEPLOYMENT"] = STUDENT_DEPLOYMENT
def set_teacher_deployment(deployment: str) -> None:
global TEACHER_DEPLOYMENT
TEACHER_DEPLOYMENT = deployment or default_model_for_backend("claude")
os.environ["TEACHER_DEPLOYMENT"] = TEACHER_DEPLOYMENT
+664
View File
@@ -0,0 +1,664 @@
"""Codex CLI backend for ReflACT."""
from __future__ import annotations
import base64
import json
import mimetypes
import os
import subprocess
import tempfile
import time
import uuid
from typing import Any
from urllib.parse import unquote, urlparse
from skillopt.model.common import (
CompatAssistantMessage,
CompatToolCall,
CompatToolFunction,
tracker,
)
CODEX_BIN = os.environ.get("CODEX_CLI_BIN", "codex")
CODEX_PROFILE = os.environ.get("CODEX_PROFILE", "review")
CODEX_SANDBOX_MODE = os.environ.get("CODEX_SANDBOX_MODE", "read-only")
TEACHER_DEPLOYMENT = os.environ.get("TEACHER_DEPLOYMENT", "gpt-5.5")
STUDENT_DEPLOYMENT = os.environ.get("STUDENT_DEPLOYMENT", "gpt-5.5")
REASONING_EFFORT: str | None = None
def _default_working_directory() -> str:
return os.environ.get("CODEX_WORKING_DIRECTORY", os.getcwd())
def _parse_data_uri(url: str) -> tuple[bytes, str]:
header, data = url.split(",", 1)
mime = header[5:].split(";", 1)[0] or "image/png"
return base64.b64decode(data), mime
def _content_to_text(
content: Any,
attachments: list[dict[str, Any]],
*,
image_counter: int,
) -> tuple[str, int]:
if isinstance(content, str):
return content, image_counter
if not isinstance(content, list):
return str(content), image_counter
parts: list[str] = []
for item in content:
if not isinstance(item, dict):
continue
item_type = item.get("type")
if item_type == "text":
parts.append(str(item.get("text", "")))
continue
if item_type != "image_url":
continue
image_counter += 1
label = f"[Attached image {image_counter}]"
parts.append(label)
image_url = item.get("image_url", {}) or {}
url = str(image_url.get("url", "") or "")
if not url:
continue
if url.startswith("data:") and ";base64," in url:
data, mime = _parse_data_uri(url)
attachments.append({"bytes": data, "mime": mime})
continue
if url.startswith("file://"):
parsed = urlparse(url)
path = unquote(parsed.path)
if path:
attachments.append({"path": path})
continue
if os.path.exists(url):
attachments.append({"path": url})
return "".join(parts), image_counter
def _simplify_tool_schemas(tools: list[dict[str, Any]] | None) -> list[dict[str, Any]]:
simplified: list[dict[str, Any]] = []
for tool in tools or []:
function = tool.get("function", tool)
simplified.append(
{
"name": function.get("name", ""),
"description": function.get("description", ""),
"parameters": function.get("parameters", {}),
}
)
return simplified
def _build_prompt_from_messages(
messages: list[dict[str, Any]],
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
structured_output: bool = False,
) -> tuple[str, list[dict[str, Any]]]:
system_parts: list[str] = []
history_parts: list[str] = []
attachments: list[dict[str, Any]] = []
image_counter = 0
def _history_line(label: str, body: str) -> str:
stripped = body.strip()
if not stripped:
return f"- {label}:"
indented = stripped.replace("\n", "\n ")
return f"- {label}: {indented}"
for message in messages:
role = str(message.get("role", "user"))
text, image_counter = _content_to_text(
message.get("content", ""),
attachments,
image_counter=image_counter,
)
if role == "system":
if text.strip():
system_parts.append(text.strip())
continue
if role == "assistant":
block = _history_line("Assistant", text)
tool_calls = message.get("tool_calls") or []
if tool_calls:
simplified_calls = []
for tool_call in tool_calls:
function = tool_call.get("function", {}) or {}
simplified_calls.append(
{
"name": function.get("name", ""),
"arguments": function.get("arguments", "{}"),
}
)
block += (
"\n Compatibility tool requests:\n"
+ json.dumps(simplified_calls, ensure_ascii=False, indent=2)
)
history_parts.append(block)
continue
if role == "tool":
tool_call_id = str(message.get("tool_call_id", "") or "")
label = f"Tool result (tool_call_id={tool_call_id})"
history_parts.append(_history_line(label, text))
continue
history_parts.append(_history_line(role.capitalize(), text))
prompt_parts: list[str] = []
system_text = "\n\n".join(part for part in system_parts if part).strip()
if system_text:
prompt_parts.append(system_text)
if tools:
simplified_tools = _simplify_tool_schemas(tools)
prompt_parts.append(
"Available compatibility tools:\n"
+ json.dumps(simplified_tools, ensure_ascii=False, indent=2)
)
prompt_parts.append(
"Do not execute these tools yourself. If you need one, request it in "
"`tool_calls`. Each `arguments` field must be a JSON string."
)
if tool_choice == "required":
prompt_parts.append(
"Tool choice policy: you must request at least one compatibility tool."
)
elif isinstance(tool_choice, dict) and tool_choice.get("type") == "function":
function = tool_choice.get("function", {}) or {}
prompt_parts.append(
"Tool choice policy: you must request the compatibility tool "
f"`{function.get('name', '')}`."
)
history_text = "\n".join(part for part in history_parts if part).strip()
if history_text:
prompt_parts.append("History:\n" + history_text)
if structured_output:
prompt_parts.append("Return only JSON matching the provided schema.")
if tools:
prompt_parts.append(
"Set `content` to the assistant-visible reply. Set `tool_calls` to "
"an empty array when no tool is needed."
)
else:
prompt_parts.append("Answer the latest user request.")
return "\n\n".join(prompt_parts), attachments
def _assistant_message_schema() -> dict[str, Any]:
return {
"type": "object",
"properties": {
"content": {"type": "string"},
"tool_calls": {
"type": "array",
"items": {
"type": "object",
"properties": {
"name": {"type": "string"},
"arguments": {"type": "string"},
},
"required": ["name", "arguments"],
"additionalProperties": False,
},
},
},
"required": ["content", "tool_calls"],
"additionalProperties": False,
}
def _materialize_attachments(
attachments: list[dict[str, Any]],
temp_dir: str,
) -> list[str]:
image_paths: list[str] = []
for index, attachment in enumerate(attachments, 1):
path = attachment.get("path")
if path:
image_paths.append(str(path))
continue
mime = str(attachment.get("mime", "image/png"))
suffix = mimetypes.guess_extension(mime) or ".png"
image_path = os.path.join(temp_dir, f"image_{index}{suffix}")
with open(image_path, "wb") as f:
f.write(attachment.get("bytes", b""))
image_paths.append(image_path)
return image_paths
def _usage_from_event(usage: dict[str, Any] | None) -> dict[str, int]:
usage = usage or {}
prompt_tokens = int(usage.get("input_tokens", 0) or 0)
completion_tokens = int(usage.get("output_tokens", 0) or 0)
return {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
def _extract_error(stdout: str, stderr: str) -> str:
for raw_line in reversed(stdout.splitlines()):
line = raw_line.strip()
if not line:
continue
try:
payload = json.loads(line)
except json.JSONDecodeError:
continue
if payload.get("type") == "turn.failed":
error = payload.get("error", {}) or {}
return str(error.get("message", "") or "Codex turn failed")
if payload.get("type") == "error":
return str(payload.get("message", "") or "Codex execution failed")
return stderr.strip() or stdout.strip() or "Codex execution failed"
def _run_codex_exec(
*,
model: str,
prompt: str,
attachments: list[dict[str, Any]],
output_schema: dict[str, Any] | None,
timeout: int | None,
) -> tuple[str, dict[str, int]]:
with tempfile.TemporaryDirectory(prefix="skillopt_codex_") as temp_dir:
output_path = os.path.join(temp_dir, "last_message.txt")
image_paths = _materialize_attachments(attachments, temp_dir)
command = [
CODEX_BIN,
"exec",
"--json",
"--ephemeral",
"--profile",
CODEX_PROFILE,
"-c",
"approval_policy=\"never\"",
"--sandbox",
CODEX_SANDBOX_MODE,
"--skip-git-repo-check",
"--cd",
_default_working_directory(),
"--model",
model,
"--output-last-message",
output_path,
]
if REASONING_EFFORT:
command.extend(["-c", f"model_reasoning_effort={json.dumps(REASONING_EFFORT)}"])
schema_path = None
if output_schema is not None:
schema_path = os.path.join(temp_dir, "schema.json")
with open(schema_path, "w", encoding="utf-8") as f:
json.dump(output_schema, f, ensure_ascii=False)
command.extend(["--output-schema", schema_path])
for image_path in image_paths:
command.extend(["--image", image_path])
command.append("-")
proc = subprocess.run(
command,
input=prompt,
text=True,
capture_output=True,
timeout=timeout,
check=False,
)
usage_info = {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
fallback_text = ""
for raw_line in proc.stdout.splitlines():
line = raw_line.strip()
if not line:
continue
try:
payload = json.loads(line)
except json.JSONDecodeError:
continue
if payload.get("type") == "item.completed":
item = payload.get("item", {}) or {}
if item.get("type") == "agent_message":
fallback_text = str(item.get("text", "") or fallback_text)
if payload.get("type") == "turn.completed":
usage_info = _usage_from_event(payload.get("usage"))
last_message = ""
if os.path.exists(output_path):
with open(output_path, encoding="utf-8") as f:
last_message = f.read().strip()
if not last_message:
last_message = fallback_text.strip()
if proc.returncode != 0:
raise RuntimeError(_extract_error(proc.stdout, proc.stderr))
if not last_message:
raise RuntimeError("Codex returned an empty final message")
return last_message, usage_info
def _tool_name_from_choice(tool_choice: str | dict[str, Any] | None) -> str | None:
if not isinstance(tool_choice, dict):
return None
if tool_choice.get("type") != "function":
return None
function = tool_choice.get("function", {}) or {}
return str(function.get("name", "") or "") or None
def _compat_message_from_payload(
payload: dict[str, Any],
*,
tool_choice: str | dict[str, Any] | None = None,
) -> CompatAssistantMessage:
content = str(payload.get("content", "") or "")
tool_calls: list[CompatToolCall] = []
for index, raw_tool_call in enumerate(payload.get("tool_calls", []) or [], 1):
if not isinstance(raw_tool_call, dict):
continue
name = str(raw_tool_call.get("name", "") or "")
arguments = raw_tool_call.get("arguments", "{}")
if not isinstance(arguments, str):
arguments = json.dumps(arguments, ensure_ascii=False)
tool_calls.append(
CompatToolCall(
id=f"tool_{index}_{uuid.uuid4().hex[:12]}",
function=CompatToolFunction(name=name, arguments=arguments),
)
)
if tool_choice == "required" and not tool_calls:
raise RuntimeError("Codex response did not request a tool under tool_choice='required'")
required_name = _tool_name_from_choice(tool_choice)
if required_name and all(
tool_call.function.name != required_name for tool_call in tool_calls
):
raise RuntimeError(
f"Codex response did not request the required tool {required_name!r}"
)
return CompatAssistantMessage(content=content, tool_calls=tool_calls)
def _chat_messages_impl(
model: str,
messages: list[dict[str, Any]],
max_completion_tokens: int,
retries: int,
stage: str,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
del max_completion_tokens
last_err = None
structured_output = bool(tools) or return_message
for attempt in range(retries):
try:
prompt, attachments = _build_prompt_from_messages(
messages,
tools=tools,
tool_choice=tool_choice,
structured_output=structured_output,
)
raw_text, usage_info = _run_codex_exec(
model=model,
prompt=prompt,
attachments=attachments,
output_schema=_assistant_message_schema() if structured_output else None,
timeout=timeout,
)
tracker.record(
stage,
usage_info["prompt_tokens"],
usage_info["completion_tokens"],
)
if not structured_output:
return raw_text, usage_info
payload = json.loads(raw_text)
compat = _compat_message_from_payload(payload, tool_choice=tool_choice)
return (compat if return_message else compat.content), usage_info
except subprocess.TimeoutExpired as exc:
last_err = RuntimeError(f"Codex CLI timed out after {timeout}s") if timeout else exc
except Exception as exc: # noqa: BLE001
last_err = exc
time.sleep(min(2 ** attempt, 30))
raise RuntimeError(f"Codex call failed after {retries} retries: {last_err}")
def chat_with_model(
model: str,
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
messages = [
{"role": "system", "content": system},
{"role": "user", "content": user},
]
return _chat_messages_impl(
model,
messages,
max_completion_tokens,
retries,
stage,
timeout=timeout,
)
def chat_messages_with_model(
model: str,
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _chat_messages_impl(
model,
messages,
max_completion_tokens,
retries,
stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_teacher(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return chat_with_model(
model=TEACHER_DEPLOYMENT,
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_with_deployment(
deployment: str,
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return chat_with_model(
model=deployment,
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_student(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return chat_with_model(
model=STUDENT_DEPLOYMENT,
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_teacher_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _chat_messages_impl(
TEACHER_DEPLOYMENT,
messages,
max_completion_tokens,
retries,
stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_messages_with_deployment(
deployment: str,
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _chat_messages_impl(
deployment,
messages,
max_completion_tokens,
retries,
stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_student_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _chat_messages_impl(
STUDENT_DEPLOYMENT,
messages,
max_completion_tokens,
retries,
stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def get_token_summary() -> dict[str, dict[str, int]]:
return tracker.summary()
def reset_token_tracker() -> None:
tracker.reset()
def set_student_deployment(deployment: str) -> None:
global STUDENT_DEPLOYMENT
STUDENT_DEPLOYMENT = deployment
os.environ["STUDENT_DEPLOYMENT"] = deployment
def set_reasoning_effort(effort: str | None) -> None:
global REASONING_EFFORT
REASONING_EFFORT = effort if effort else None
def set_teacher_deployment(deployment: str) -> None:
global TEACHER_DEPLOYMENT
TEACHER_DEPLOYMENT = deployment
os.environ["TEACHER_DEPLOYMENT"] = deployment
File diff suppressed because it is too large Load Diff
+226
View File
@@ -0,0 +1,226 @@
"""Shared model utilities for ReflACT backends."""
from __future__ import annotations
import json
import threading
from dataclasses import dataclass, field
from typing import Any
_RESPONSES_API_MODELS = {
"gpt-5.3-codex",
"gpt-5.1-codex",
"gpt-5.2-codex",
"gpt-5-codex",
"codex-mini",
"gpt-5.4-pro",
}
_BACKEND_DEFAULT_MODELS = {
"azure_openai": "gpt-5.5",
"openai_chat": "gpt-5.5",
"codex": "gpt-5.5",
"codex_exec": "gpt-5.5",
"claude": "claude-sonnet-4-6",
"claude_chat": "claude-sonnet-4-6",
"claude_code_exec": "claude-sonnet-4-6",
"qwen_chat": "Qwen/Qwen3.5-4B",
}
_BACKEND_ALIASES = {
"azure": "azure_openai",
"azure_openai": "azure_openai",
"azure-openai": "azure_openai",
"openai_chat": "openai_chat",
"openai": "codex",
"codex": "codex",
"codex_exec": "codex_exec",
"claude": "claude_chat",
"claude_chat": "claude_chat",
"claude_code_exec": "claude_code_exec",
"anthropic": "claude_chat",
"qwen": "qwen_chat",
"qwen_chat": "qwen_chat",
}
def normalize_backend_name(name: str | None) -> str:
normalized = str(name or "").strip().lower()
return _BACKEND_ALIASES.get(normalized, normalized or "azure_openai")
def default_model_for_backend(backend: str | None) -> str:
return _BACKEND_DEFAULT_MODELS.get(
normalize_backend_name(backend),
_BACKEND_DEFAULT_MODELS["azure_openai"],
)
def needs_responses_api(model: str) -> bool:
normalized = str(model or "").strip().lower()
return any(
normalized == prefix or normalized.startswith(prefix + "-")
for prefix in _RESPONSES_API_MODELS
)
class TokenTracker:
def __init__(self) -> None:
self._lock = threading.Lock()
self._data: dict[str, dict[str, int]] = {}
def record(self, stage: str, prompt_tokens: int, completion_tokens: int) -> None:
with self._lock:
if stage not in self._data:
self._data[stage] = {
"calls": 0,
"prompt_tokens": 0,
"completion_tokens": 0,
}
entry = self._data[stage]
entry["calls"] += 1
entry["prompt_tokens"] += prompt_tokens
entry["completion_tokens"] += completion_tokens
def summary(self) -> dict[str, dict[str, int]]:
with self._lock:
out: dict[str, dict[str, int]] = {}
total_prompt = total_completion = total_calls = 0
for stage, entry in sorted(self._data.items()):
prompt_tokens = entry["prompt_tokens"]
completion_tokens = entry["completion_tokens"]
out[stage] = {
"calls": entry["calls"],
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
total_prompt += prompt_tokens
total_completion += completion_tokens
total_calls += entry["calls"]
out["_total"] = {
"calls": total_calls,
"prompt_tokens": total_prompt,
"completion_tokens": total_completion,
"total_tokens": total_prompt + total_completion,
}
return out
def reset(self) -> None:
with self._lock:
self._data.clear()
tracker = TokenTracker()
@dataclass
class CompatToolFunction:
name: str
arguments: str
def model_dump(self, mode: str = "json") -> dict[str, str]:
del mode
return {
"name": self.name,
"arguments": self.arguments,
}
@dataclass
class CompatToolCall:
id: str
function: CompatToolFunction
type: str = "function"
def model_dump(self, mode: str = "json") -> dict[str, Any]:
del mode
return {
"id": self.id,
"type": self.type,
"function": self.function.model_dump(),
}
@dataclass
class CompatAssistantMessage:
content: str
tool_calls: list[CompatToolCall] = field(default_factory=list)
metadata: dict[str, Any] = field(default_factory=dict)
def model_dump(self, mode: str = "json") -> dict[str, Any]:
del mode
data: dict[str, Any] = {"role": "assistant", "content": self.content}
if self.tool_calls:
data["tool_calls"] = [tool_call.model_dump() for tool_call in self.tool_calls]
return data
def usage_from_openai_usage(usage: Any) -> dict[str, int]:
if not usage:
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
prompt_tokens = getattr(usage, "prompt_tokens", 0) or 0
completion_tokens = getattr(usage, "completion_tokens", 0) or 0
total_tokens = getattr(usage, "total_tokens", 0) or (prompt_tokens + completion_tokens)
return {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens,
}
def usage_from_responses_usage(usage: Any) -> dict[str, int]:
if not usage:
return {"prompt_tokens": 0, "completion_tokens": 0, "total_tokens": 0}
prompt_tokens = getattr(usage, "input_tokens", 0) or 0
completion_tokens = getattr(usage, "output_tokens", 0) or 0
return {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": prompt_tokens + completion_tokens,
}
def compat_message_from_chat_message(message: Any) -> CompatAssistantMessage:
content = getattr(message, "content", "") or ""
tool_calls = []
for tool_call in getattr(message, "tool_calls", None) or []:
function = getattr(tool_call, "function", None)
tool_calls.append(
CompatToolCall(
id=getattr(tool_call, "id", "") or "",
function=CompatToolFunction(
name=getattr(function, "name", "") or "",
arguments=getattr(function, "arguments", "") or "{}",
),
)
)
return CompatAssistantMessage(content=content, tool_calls=tool_calls)
def compat_message_from_responses_output(output: list[Any]) -> CompatAssistantMessage:
text_parts: list[str] = []
tool_calls: list[CompatToolCall] = []
for item in output:
item_type = getattr(item, "type", "") or ""
if item_type == "function_call":
raw_arguments = getattr(item, "arguments", None)
if raw_arguments is None:
raw_arguments = json.dumps(getattr(item, "input", {}) or {})
tool_calls.append(
CompatToolCall(
id=getattr(item, "call_id", "") or getattr(item, "id", "") or "",
function=CompatToolFunction(
name=getattr(item, "name", "") or "",
arguments=str(raw_arguments or "{}"),
),
)
)
continue
if item_type != "message":
continue
for part in getattr(item, "content", []) or []:
part_type = getattr(part, "type", "") or ""
if part_type in {"output_text", "text"}:
text_parts.append(getattr(part, "text", "") or "")
return CompatAssistantMessage(content="".join(text_parts), tool_calls=tool_calls)
+277
View File
@@ -0,0 +1,277 @@
"""OpenAI-compatible Qwen chat backend for the student path."""
from __future__ import annotations
import json
import os
import threading
import time
import urllib.error
import urllib.request
from typing import Any
from skillopt.model.common import (
CompatAssistantMessage,
CompatToolCall,
CompatToolFunction,
TokenTracker,
default_model_for_backend,
)
BASE_URL = os.environ.get("QWEN_CHAT_BASE_URL", "http://localhost:8000/v1")
API_KEY = os.environ.get("QWEN_CHAT_API_KEY", "")
TIMEOUT_SECONDS = float(os.environ.get("QWEN_CHAT_TIMEOUT_SECONDS", "300") or 300)
MAX_TOKENS = int(os.environ.get("QWEN_CHAT_MAX_TOKENS", "8000") or 8000)
TEMPERATURE: float | None = None
_raw_temperature = os.environ.get("QWEN_CHAT_TEMPERATURE", "0.7").strip()
if _raw_temperature:
TEMPERATURE = float(_raw_temperature)
ENABLE_THINKING = os.environ.get("QWEN_CHAT_ENABLE_THINKING", "false").strip().lower() in {
"1",
"true",
"yes",
"on",
}
STUDENT_DEPLOYMENT = os.environ.get(
"STUDENT_DEPLOYMENT",
default_model_for_backend("qwen_chat"),
)
_config_lock = threading.Lock()
tracker = TokenTracker()
def _chat_url() -> str:
base = BASE_URL.rstrip("/")
if base.endswith("/chat/completions"):
return base
return f"{base}/chat/completions"
def _json_safe(value: Any) -> Any:
if value is None or isinstance(value, (str, int, float, bool)):
return value
if isinstance(value, list):
return [_json_safe(item) for item in value]
if isinstance(value, dict):
return {str(key): _json_safe(val) for key, val in value.items()}
model_dump = getattr(value, "model_dump", None)
if callable(model_dump):
try:
return model_dump(mode="json")
except TypeError:
return model_dump()
return str(value)
def _usage_from_payload(payload: dict[str, Any]) -> dict[str, int]:
usage = payload.get("usage") or {}
prompt_tokens = int(usage.get("prompt_tokens") or usage.get("input_tokens") or 0)
completion_tokens = int(usage.get("completion_tokens") or usage.get("output_tokens") or 0)
total_tokens = int(usage.get("total_tokens") or (prompt_tokens + completion_tokens))
return {
"prompt_tokens": prompt_tokens,
"completion_tokens": completion_tokens,
"total_tokens": total_tokens,
}
def _compat_message_from_payload(message: dict[str, Any], choice: dict[str, Any]) -> CompatAssistantMessage:
content = message.get("content") or ""
if not isinstance(content, str):
content = json.dumps(content, ensure_ascii=False)
tool_calls: list[CompatToolCall] = []
for index, tool_call in enumerate(message.get("tool_calls") or [], start=1):
function = tool_call.get("function") or {}
tool_calls.append(
CompatToolCall(
id=str(tool_call.get("id") or f"qwen_tool_{index}"),
type=str(tool_call.get("type") or "function"),
function=CompatToolFunction(
name=str(function.get("name") or ""),
arguments=str(function.get("arguments") or "{}"),
),
)
)
return CompatAssistantMessage(
content=content,
tool_calls=tool_calls,
metadata={
"finish_reason": choice.get("finish_reason"),
"choice0": _json_safe(choice),
},
)
def _post_chat_completion(payload: dict[str, Any], timeout: float | None) -> dict[str, Any]:
headers = {"Content-Type": "application/json"}
if API_KEY:
headers["Authorization"] = f"Bearer {API_KEY}"
req = urllib.request.Request(
_chat_url(),
data=json.dumps(payload, ensure_ascii=False).encode("utf-8"),
headers=headers,
method="POST",
)
try:
with urllib.request.urlopen(req, timeout=timeout or TIMEOUT_SECONDS) as resp:
raw = resp.read().decode("utf-8")
except urllib.error.HTTPError as e:
body = e.read().decode("utf-8", errors="replace")
raise RuntimeError(f"Qwen chat API returned HTTP {e.code}: {body}") from e
except urllib.error.URLError as e:
raise RuntimeError(f"Qwen chat API request failed: {e}") from e
try:
return json.loads(raw)
except json.JSONDecodeError as e:
raise RuntimeError(f"Qwen chat API returned non-JSON response: {raw[:1000]}") from e
def _chat_messages_impl(
messages: list[dict[str, Any]],
max_completion_tokens: int,
retries: int,
stage: str,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
deployment: str | None = None,
timeout: float | None = None,
) -> tuple[Any, dict[str, int]]:
payload: dict[str, Any] = {
"model": deployment or STUDENT_DEPLOYMENT,
"messages": _json_safe(messages),
"max_tokens": min(max_completion_tokens, MAX_TOKENS),
}
payload["chat_template_kwargs"] = {"enable_thinking": ENABLE_THINKING}
if TEMPERATURE is not None:
payload["temperature"] = TEMPERATURE
if tools:
payload["tools"] = _json_safe(tools)
if tool_choice is not None:
payload["tool_choice"] = _json_safe(tool_choice)
last_err: Exception | None = None
for attempt in range(retries):
try:
data = _post_chat_completion(payload, timeout)
choices = data.get("choices") or []
if not choices:
raise RuntimeError(f"Qwen chat API returned no choices: {data}")
choice0 = choices[0]
message = choice0.get("message") or {}
text = message.get("content") or ""
if not isinstance(text, str):
text = json.dumps(text, ensure_ascii=False)
usage_info = _usage_from_payload(data)
tracker.record(stage, usage_info["prompt_tokens"], usage_info["completion_tokens"])
if return_message:
return _compat_message_from_payload(message, choice0), usage_info
return text, usage_info
except Exception as e: # noqa: BLE001
last_err = e
time.sleep(min(2 ** attempt, 30))
raise RuntimeError(f"Qwen chat call failed after {retries} retries: {last_err}")
def configure_qwen_chat(
*,
base_url: str | None = None,
api_key: str | None = None,
temperature: float | str | None = None,
timeout_seconds: float | str | None = None,
max_tokens: int | str | None = None,
enable_thinking: bool | str | None = None,
) -> None:
global BASE_URL, API_KEY, TEMPERATURE, TIMEOUT_SECONDS, MAX_TOKENS, ENABLE_THINKING
with _config_lock:
if base_url is not None:
BASE_URL = str(base_url).strip() or BASE_URL
os.environ["QWEN_CHAT_BASE_URL"] = BASE_URL
if api_key is not None:
API_KEY = str(api_key).strip()
os.environ["QWEN_CHAT_API_KEY"] = API_KEY
if temperature is not None:
raw = str(temperature).strip()
TEMPERATURE = float(raw) if raw else None
os.environ["QWEN_CHAT_TEMPERATURE"] = raw
if timeout_seconds is not None:
TIMEOUT_SECONDS = float(timeout_seconds)
os.environ["QWEN_CHAT_TIMEOUT_SECONDS"] = str(timeout_seconds)
if max_tokens is not None:
MAX_TOKENS = int(max_tokens)
os.environ["QWEN_CHAT_MAX_TOKENS"] = str(max_tokens)
if enable_thinking is not None:
if isinstance(enable_thinking, str):
ENABLE_THINKING = enable_thinking.strip().lower() in {"1", "true", "yes", "on"}
else:
ENABLE_THINKING = bool(enable_thinking)
os.environ["QWEN_CHAT_ENABLE_THINKING"] = "true" if ENABLE_THINKING else "false"
def get_max_tokens() -> int:
return MAX_TOKENS
def chat_student(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
timeout: float | None = None,
) -> tuple[str, dict[str, int]]:
del reasoning_effort
messages = [{"role": "system", "content": system}, {"role": "user", "content": user}]
return _chat_messages_impl(
messages,
max_completion_tokens,
retries,
stage,
timeout=timeout,
)
def chat_student_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
reasoning_effort: str | None = None,
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: float | None = None,
) -> tuple[Any, dict[str, int]]:
del reasoning_effort
return _chat_messages_impl(
messages,
max_completion_tokens,
retries,
stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def get_token_summary() -> dict[str, dict[str, int]]:
return tracker.summary()
def reset_token_tracker() -> None:
tracker.reset()
def set_reasoning_effort(effort: str | None) -> None:
del effort
def set_student_deployment(deployment: str) -> None:
global STUDENT_DEPLOYMENT
STUDENT_DEPLOYMENT = deployment or default_model_for_backend("qwen_chat")
os.environ["STUDENT_DEPLOYMENT"] = STUDENT_DEPLOYMENT
+236
View File
@@ -0,0 +1,236 @@
"""Runtime backend router for ReflACT model calls."""
from __future__ import annotations
import os
from typing import Any
from . import azure_openai, claude_backend, codex_backend
from .common import normalize_backend_name
_ACTIVE_BACKEND = normalize_backend_name(
os.environ.get("REFLACT_MODEL_BACKEND", "azure_openai")
)
def _backend_module(name: str):
if name == "azure_openai":
return azure_openai
if name == "codex":
return codex_backend
if name == "claude":
return claude_backend
raise ValueError(f"Unknown backend: {name!r}")
def _all_backend_modules() -> list[Any]:
return [azure_openai, codex_backend, claude_backend]
def set_backend(name: str | None) -> str:
"""Select the active model backend for subsequent calls."""
global _ACTIVE_BACKEND
normalized = normalize_backend_name(name)
if normalized not in {"azure_openai", "codex", "claude"}:
valid = ", ".join(sorted({"azure_openai", "codex", "claude"}))
raise ValueError(f"Unknown backend {name!r}. Expected one of: {valid}")
_ACTIVE_BACKEND = normalized
os.environ["REFLACT_MODEL_BACKEND"] = normalized
return _ACTIVE_BACKEND
def get_backend_name() -> str:
return _ACTIVE_BACKEND
def chat_teacher(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_teacher(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_student(
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_student(
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_with_deployment(
deployment: str,
system: str,
user: str,
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
timeout: int | None = None,
) -> tuple[str, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_with_deployment(
deployment=deployment,
system=system,
user=user,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
timeout=timeout,
)
def chat_teacher_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "teacher",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_teacher_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_student_messages(
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "student",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_student_messages(
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def chat_messages_with_deployment(
deployment: str,
messages: list[dict[str, Any]],
max_completion_tokens: int = 16384,
retries: int = 5,
stage: str = "custom",
*,
tools: list[dict[str, Any]] | None = None,
tool_choice: str | dict[str, Any] | None = None,
return_message: bool = False,
timeout: int | None = None,
) -> tuple[Any, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).chat_messages_with_deployment(
deployment=deployment,
messages=messages,
max_completion_tokens=max_completion_tokens,
retries=retries,
stage=stage,
tools=tools,
tool_choice=tool_choice,
return_message=return_message,
timeout=timeout,
)
def get_token_summary() -> dict[str, dict[str, int]]:
return _backend_module(_ACTIVE_BACKEND).get_token_summary()
def reset_token_tracker() -> None:
_backend_module(_ACTIVE_BACKEND).reset_token_tracker()
def set_reasoning_effort(effort: str | None) -> None:
for module in _all_backend_modules():
module.set_reasoning_effort(effort)
def set_student_deployment(deployment: str) -> None:
for module in _all_backend_modules():
module.set_student_deployment(deployment)
def set_teacher_deployment(deployment: str) -> None:
for module in _all_backend_modules():
module.set_teacher_deployment(deployment)
def configure_azure_openai(
*,
endpoint: str | None = None,
api_version: str | None = None,
api_key: str | None = None,
auth_mode: str | None = None,
ad_scope: str | None = None,
managed_identity_client_id: str | None = None,
teacher_endpoint: str | None = None,
teacher_api_version: str | None = None,
teacher_api_key: str | None = None,
teacher_auth_mode: str | None = None,
teacher_ad_scope: str | None = None,
teacher_managed_identity_client_id: str | None = None,
student_endpoint: str | None = None,
student_api_version: str | None = None,
student_api_key: str | None = None,
student_auth_mode: str | None = None,
student_ad_scope: str | None = None,
student_managed_identity_client_id: str | None = None,
) -> None:
azure_openai.configure_azure_openai(
endpoint=endpoint,
api_version=api_version,
api_key=api_key,
auth_mode=auth_mode,
ad_scope=ad_scope,
managed_identity_client_id=managed_identity_client_id,
teacher_endpoint=teacher_endpoint,
teacher_api_version=teacher_api_version,
teacher_api_key=teacher_api_key,
teacher_auth_mode=teacher_auth_mode,
teacher_ad_scope=teacher_ad_scope,
teacher_managed_identity_client_id=teacher_managed_identity_client_id,
student_endpoint=student_endpoint,
student_api_version=student_api_version,
student_api_key=student_api_key,
student_auth_mode=student_auth_mode,
student_ad_scope=student_ad_scope,
student_managed_identity_client_id=student_managed_identity_client_id,
)