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:
@@ -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)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
}
|
||||
@@ -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
|
||||
@@ -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
@@ -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)
|
||||
@@ -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
|
||||
@@ -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,
|
||||
)
|
||||
Reference in New Issue
Block a user