Files
liqiang b119135836
Build latest book artifacts / build (push) Canceled after 0s
dependency resolution / resolve (3.11) (push) Canceled after 0s
dependency resolution / resolve (3.13) (push) Canceled after 0s
deploy-pages / build (push) Canceled after 0s
deploy-pages / deploy (push) Canceled after 0s
i18n consistency check / check (push) Canceled after 0s
provider adoption tests / test (chapter2/context-compression) (push) Canceled after 0s
provider adoption tests / test (chapter2/prompt-injection) (push) Canceled after 0s
provider adoption tests / test (chapter2/system-hint) (push) Canceled after 0s
provider adoption tests / test (chapter3/log-sanitization) (push) Canceled after 0s
web-search-agent tests / test (push) Canceled after 0s
web-search-agent tests / agentbook (push) Canceled after 0s
ai-agent-book 精选快照(<2MB 代码与文档,来自 github.com/bojieli/ai-agent-book)
2026-08-20 13:12:50 +00:00

119 lines
4.3 KiB
Python
Raw Permalink Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
"""OpenAI Responses / 兼容 Chat Completions 客户端:统一证据回执。
约定与 chapter8/self-modifying-agent/llm_generator.py 一致:每次真实调用返回
(content, receipt)receipt 含原始请求、原始响应、Token 用量、延迟与
请求/响应哈希,不记录凭据值。凭证从环境变量读取,支持 ark / openrouter / openai。
"""
from __future__ import annotations
import hashlib
import json
import os
import time
from typing import Any, Dict, List, Tuple
_PROVIDERS = {
"openrouter": ("OPENROUTER_API_KEY", "https://openrouter.ai/api/v1"),
"ark": ("ARK_API_KEY", "https://ark.cn-beijing.volces.com/api/v3"),
"openai": ("OPENAI_API_KEY", None),
}
_DEFAULT_MODELS = {
"openrouter": "openai/gpt-5.6-sol",
"ark": "doubao-seed-1-6-250615",
"openai": "gpt-5.6-sol",
}
def make_client(provider: str) -> Tuple[Any, Dict[str, Any]]:
# openai 包只在真实路径才需要,惰性导入保证离线路径零依赖。
from openai import OpenAI
if provider not in _PROVIDERS:
raise ValueError(f"不支持的 provider{provider}(可选:{sorted(_PROVIDERS)}")
env_name, base_url = _PROVIDERS[provider]
key = os.getenv(env_name)
if not key:
raise RuntimeError(f"真实 LLM 路径需要设置环境变量 {env_name}")
client = OpenAI(api_key=key, base_url=base_url) if base_url else OpenAI(api_key=key)
api = "responses" if provider in {"openai", "openrouter"} else "chat/completions"
backend = {
"provider": provider,
"endpoint": (base_url or "https://api.openai.com/v1") + f"/{api}",
"credential_env": env_name,
}
return client, backend
def default_model(provider: str) -> str:
if provider == "ark":
return os.getenv("ARK_MODEL", _DEFAULT_MODELS["ark"])
return _DEFAULT_MODELS[provider]
def chat(
messages: List[Dict[str, str]],
*,
provider: str,
model: str | None = None,
seed: int = 8901,
max_tokens: int = 5000,
) -> Tuple[str, Dict[str, Any]]:
"""发起一次结构化 JSON 调用,返回 ``(文本内容, 证据回执)``。"""
client, backend = make_client(provider)
selected = model or default_model(provider)
started = time.perf_counter()
if provider in {"openai", "openrouter"}:
request = {
"model": selected,
"input": messages,
"reasoning": {"effort": "medium"},
"max_output_tokens": max_tokens,
"text": {"format": {"type": "json_object"}},
"store": False,
}
response = client.responses.create(**request)
content = response.output_text
else:
request = {
"model": selected,
"messages": messages,
"temperature": 0,
"seed": seed,
"max_tokens": max_tokens,
"response_format": {"type": "json_object"},
}
response = client.chat.completions.create(**request)
content = response.choices[0].message.content or ""
elapsed = time.perf_counter() - started
raw = response.model_dump(mode="json", exclude_none=True)
usage = raw.get("usage") or {}
cost = usage.get("cost")
prompt_tokens = usage.get("input_tokens", usage.get("prompt_tokens", 0))
completion_tokens = usage.get("output_tokens", usage.get("completion_tokens", 0))
receipt = {
"backend": {**backend, "model": selected, "credential_value_recorded": False},
"request": request,
"response": raw,
"request_sha256": hashlib.sha256(
json.dumps(request, sort_keys=True).encode()
).hexdigest(),
"response_sha256": hashlib.sha256(
json.dumps(raw, sort_keys=True).encode()
).hexdigest(),
"elapsed_seconds": round(elapsed, 6),
"usage": {
"prompt_tokens": int(prompt_tokens or 0),
"completion_tokens": int(completion_tokens or 0),
"total_tokens": int(usage.get("total_tokens") or 0),
"provider_reported_cost_usd": float(cost) if cost is not None else None,
"cost_qualification": (
"provider-native usage.cost"
if cost is not None
else "provider did not expose monetary cost; no price was guessed"
),
},
}
return content, receipt