diff --git a/client/README.md b/client/README.md index 5d29eda6..354d1123 100644 --- a/client/README.md +++ b/client/README.md @@ -166,7 +166,7 @@ export OPENROUTER_API_KEY=sk-or-... ### `llm` 字段说明 -- `adapter`:模型适配器,当前支持 `openai`、`openrouter`、`anthropic`。不写时默认是 `openai`。 +- `adapter`:模型适配器,当前支持 `openai`、`openrouter`、`anthropic`、`huawei-modelarts`。不写时默认是 `openai`。 - `model`:模型名,例如 `gpt-4o-mini`、`gpt-4.1`、`qwen/qwen3-coder:free`。 - `api_key`:模型服务密钥。也可以通过环境变量提供。 - `endpoint`:自定义 provider base URL。对 `openrouter`/`anthropic` 来说分别映射到各自 SDK 的 `base_url`。 @@ -283,6 +283,35 @@ export ANTHROPIC_MODEL=claude-sonnet-4-5 - `client/examples/config.anthropic.example.json` +Huawei ModelArts: + +```json +{ + "llm": { + "adapter": "huawei-modelarts", + "model": "glm-5" + } +} +``` + +配合环境变量: + +```bash +export HUAWEI_MODELARTS_API_KEY=your-key +# 可选,默认就是 https://api.modelarts-maas.com/v2 +export HUAWEI_MODELARTS_BASE_URL=https://api.modelarts-maas.com/v2 +``` + +说明: + +- 默认 base URL 是 `https://api.modelarts-maas.com/v2` +- runtime 实际调用 OpenAI SDK 风格的 `/chat/completions` +- 也可以通过 `llm.endpoint` 或 `--endpoint` 临时覆盖 base URL + +仓库示例文件: + +- `client/examples/config.huawei-modelarts.example.json` + OpenAI-compatible / 自建模型网关: ```json @@ -414,6 +443,7 @@ python -m client doctor - `openai` adapter:需要 `langchain-openai` - `openrouter` adapter:需要 `openai` - `anthropic` adapter:需要 `anthropic` +- `huawei-modelarts` adapter:需要 `openai` 否则 `doctor` 会提示 adapter probe 或依赖缺失,runtime 也无法正常启动。 diff --git a/client/commands/info.py b/client/commands/info.py index 9f8de17c..2716ff38 100644 --- a/client/commands/info.py +++ b/client/commands/info.py @@ -53,6 +53,9 @@ def build_doctor_report( "openrouter": bool(config.llm.api_key or os.getenv("OPENROUTER_API_KEY")), "openai": bool(config.llm.api_key or os.getenv("OPENAI_API_KEY")), "anthropic": bool(config.llm.api_key or os.getenv("ANTHROPIC_API_KEY")), + "huawei-modelarts": bool( + config.llm.api_key or os.getenv("HUAWEI_MODELARTS_API_KEY") + ), } workspace = Path(config.workspace_dir).expanduser().resolve() user_dir = Path(config.user_dir).expanduser().resolve() @@ -97,7 +100,7 @@ def build_doctor_report( warnings: list[str] = [] if not diagnostics["workspace_exists"]: warnings.append("workspace_dir does not exist") - if adapter not in {"openai", "openrouter", "anthropic"}: + if adapter not in {"openai", "openrouter", "anthropic", "huawei-modelarts"}: warnings.append(f"unsupported adapter configured: {adapter}") if not diagnostics["llm"]["api_key_present"]: warnings.append(f"missing API key for adapter={adapter}") @@ -115,6 +118,8 @@ def build_doctor_report( warnings.append("langchain-openai is required for openai adapter") if adapter == "anthropic" and not deps["anthropic_sdk_installed"]: warnings.append("anthropic SDK is required for anthropic adapter") + if adapter == "huawei-modelarts" and not deps["openai_sdk_installed"]: + warnings.append("openai SDK is required for huawei-modelarts adapter") diagnostics["warnings"] = warnings diagnostics["ok"] = len(warnings) == 0 return diagnostics diff --git a/client/examples/config.huawei-modelarts.example.json b/client/examples/config.huawei-modelarts.example.json new file mode 100644 index 00000000..4dbc2708 --- /dev/null +++ b/client/examples/config.huawei-modelarts.example.json @@ -0,0 +1,6 @@ +{ + "llm": { + "adapter": "huawei-modelarts", + "model": "modelarts-pro" + } +} diff --git a/client/main.py b/client/main.py index 9b3173d2..904e34d1 100644 --- a/client/main.py +++ b/client/main.py @@ -1538,7 +1538,7 @@ def _build_parser() -> argparse.ArgumentParser: parser = argparse.ArgumentParser(description="DARE external CLI") parser.add_argument("--workspace", default=str(Path.cwd()), help="workspace root path") parser.add_argument("--user-dir", default=str(Path.home()), help="user directory path") - parser.add_argument("--adapter", default=None, help="llm adapter override (openai/openrouter/anthropic)") + parser.add_argument("--adapter", default=None, help="llm adapter override (openai/openrouter/anthropic/huawei-modelarts)") parser.add_argument("--model", default=None, help="llm model override") parser.add_argument("--api-key", default=None, help="llm api key override") parser.add_argument("--endpoint", default=None, help="llm endpoint override") diff --git a/dare_framework/model/__init__.py b/dare_framework/model/__init__.py index 45e2c1da..4948cdb6 100644 --- a/dare_framework/model/__init__.py +++ b/dare_framework/model/__init__.py @@ -9,7 +9,12 @@ from dare_framework.model.builtin_prompt_loader import BuiltInPromptLoader from dare_framework.model.filesystem_prompt_loader import FileSystemPromptLoader from dare_framework.model.layered_prompt_store import LayeredPromptStore -from dare_framework.model.adapters import AnthropicModelAdapter, OpenAIModelAdapter, OpenRouterModelAdapter +from dare_framework.model.adapters import ( + AnthropicModelAdapter, + HuaweiModelArtsModelAdapter, + OpenAIModelAdapter, + OpenRouterModelAdapter, +) __all__ = [ "IModelAdapter", @@ -25,6 +30,7 @@ "FileSystemPromptLoader", "LayeredPromptStore", "AnthropicModelAdapter", + "HuaweiModelArtsModelAdapter", "OpenAIModelAdapter", "OpenRouterModelAdapter", ] diff --git a/dare_framework/model/adapters/__init__.py b/dare_framework/model/adapters/__init__.py index e0662803..99d01c6f 100644 --- a/dare_framework/model/adapters/__init__.py +++ b/dare_framework/model/adapters/__init__.py @@ -1,7 +1,13 @@ """Model adapters.""" from dare_framework.model.adapters.anthropic_adapter import AnthropicModelAdapter +from dare_framework.model.adapters.huawei_modelarts_adapter import HuaweiModelArtsModelAdapter from dare_framework.model.adapters.openai_adapter import OpenAIModelAdapter from dare_framework.model.adapters.openrouter_adapter import OpenRouterModelAdapter -__all__ = ["AnthropicModelAdapter", "OpenAIModelAdapter", "OpenRouterModelAdapter"] +__all__ = [ + "AnthropicModelAdapter", + "HuaweiModelArtsModelAdapter", + "OpenAIModelAdapter", + "OpenRouterModelAdapter", +] diff --git a/dare_framework/model/adapters/huawei_modelarts_adapter.py b/dare_framework/model/adapters/huawei_modelarts_adapter.py new file mode 100644 index 00000000..9e234269 --- /dev/null +++ b/dare_framework/model/adapters/huawei_modelarts_adapter.py @@ -0,0 +1,354 @@ +"""Huawei ModelArts model adapter using OpenAI SDK-compatible chat completions.""" + +from __future__ import annotations + +import json +import os +from typing import Any + +from dare_framework.model.kernel import IModelAdapter +from dare_framework.model.types import GenerateOptions, ModelInput, ModelResponse + +_MODELARTS_BASE_URL = "https://api.modelarts-maas.com/v2" + + +class HuaweiModelArtsModelAdapter(IModelAdapter): + """Model adapter for Huawei ModelArts MaaS (OpenAI-compatible).""" + + def __init__( + self, + *, + name: str | None = None, + api_key: str | None = None, + model: str | None = None, + base_url: str | None = None, + http_client_options: dict[str, Any] | None = None, + extra: dict[str, Any] | None = None, + ) -> None: + self._name = name or "huawei-modelarts" + self._api_key = api_key or os.getenv("HUAWEI_MODELARTS_API_KEY") + self._model = model + self._base_url = base_url or os.getenv("HUAWEI_MODELARTS_BASE_URL") or _MODELARTS_BASE_URL + self._http_client_options = dict(http_client_options or {}) + self._extra = dict(extra or {}) + self._client: Any = None + + if not self._api_key: + raise ValueError( + "Huawei ModelArts API key is required. Set HUAWEI_MODELARTS_API_KEY environment variable." + ) + if not self._model: + raise ValueError("Huawei ModelArts model is required. Set llm.model in config or pass --model.") + + @property + def name(self) -> str: + return self._name + + @property + def model(self) -> str: + return self._model or "" + + @property + def model_name(self) -> str: + return self.model + + async def generate( + self, + model_input: ModelInput, + *, + options: GenerateOptions | None = None, + ) -> ModelResponse: + client = self._ensure_client() + messages = _serialize_messages(model_input.messages) + + api_params: dict[str, Any] = { + "model": self.model, + "messages": messages, + } + if model_input.tools: + api_params["tools"] = [ + { + "type": "function", + "function": { + "name": tool.name, + "description": tool.description, + "parameters": tool.input_schema, + }, + } + for tool in model_input.tools + ] + + if self._extra: + api_params.update(self._extra) + if options is not None: + if options.temperature is not None: + api_params["temperature"] = options.temperature + if options.max_tokens is not None: + api_params["max_tokens"] = options.max_tokens + if options.top_p is not None: + api_params["top_p"] = options.top_p + if options.stop is not None: + api_params["stop"] = options.stop + + response = await client.chat.completions.create(**api_params) + message = response.choices[0].message + content = message.content or "" + tool_calls = _extract_tool_calls(message) + thinking_content = _extract_thinking_content(message) + + usage = None + if response.usage: + usage = { + "prompt_tokens": response.usage.prompt_tokens, + "completion_tokens": response.usage.completion_tokens, + "total_tokens": response.usage.total_tokens, + } + reasoning_tokens = _extract_reasoning_tokens(response) + if reasoning_tokens is not None: + if usage is None: + usage = {} + usage["reasoning_tokens"] = reasoning_tokens + + return ModelResponse( + content=content, + tool_calls=tool_calls, + usage=usage, + thinking_content=thinking_content, + metadata={ + "model": self.model, + "finish_reason": response.choices[0].finish_reason, + "base_url": self._base_url, + }, + ) + + def _ensure_client(self) -> Any: + if self._client is None: + self._client = self._build_client() + return self._client + + def _build_client(self) -> Any: + try: + from openai import AsyncOpenAI + except ImportError as exc: + raise ImportError( + "OpenAI SDK is required for Huawei ModelArts. Install with: pip install openai" + ) from exc + + client_kwargs: dict[str, Any] = { + "api_key": self._api_key, + "base_url": self._base_url, + } + + http_client = _build_async_http_client(self._http_client_options) + if http_client is not None: + client_kwargs["http_client"] = http_client + + return AsyncOpenAI(**client_kwargs) + + +def _serialize_messages(messages: list[Any]) -> list[dict[str, Any]]: + serialized: list[dict[str, Any]] = [] + for msg in messages: + role = str(getattr(msg, "role", "user")) + payload: dict[str, Any] = { + "role": role, + "content": _serialize_openai_compatible_content(msg), + } + if role == "assistant": + tool_calls = _normalize_tool_calls_for_openai_sdk(_extract_message_tool_calls(msg)) + if tool_calls: + payload["tool_calls"] = tool_calls + tool_call_id = _extract_message_tool_call_id(msg) + name = getattr(msg, "name", None) + if role == "tool" and tool_call_id: + payload["tool_call_id"] = tool_call_id + elif name: + payload["name"] = name + serialized.append(payload) + return serialized + + +def _serialize_openai_compatible_content(message: Any) -> Any: + text = _message_text(message) + attachments = list(getattr(message, "attachments", []) or []) + if not attachments: + return text + + content: list[dict[str, Any]] = [] + if text: + content.append({"type": "text", "text": text}) + for attachment in attachments: + if str(getattr(attachment, "kind", "")).strip().lower() != "image": + raise ValueError("unsupported attachment kind for Huawei ModelArts serialization") + content.append({"type": "image_url", "image_url": {"url": attachment.uri}}) + return content + + +def _message_text(message: Any) -> str: + text = getattr(message, "text", None) + if isinstance(text, str): + return text + return "" + + +def _extract_message_tool_calls(message: Any) -> Any: + data = getattr(message, "data", None) + if isinstance(data, dict) and isinstance(data.get("tool_calls"), list): + return data["tool_calls"] + return [] + + +def _extract_message_tool_call_id(message: Any) -> str | None: + data = getattr(message, "data", None) + if isinstance(data, dict): + tool_call_id = data.get("tool_call_id") + if isinstance(tool_call_id, str) and tool_call_id.strip(): + return tool_call_id + name = getattr(message, "name", None) + if isinstance(name, str) and name.strip(): + return name + return None + + +def _normalize_tool_calls_for_openai_sdk(tool_calls: Any) -> list[dict[str, Any]]: + if not isinstance(tool_calls, list): + return [] + + normalized: list[dict[str, Any]] = [] + for call in tool_calls: + if not isinstance(call, dict): + continue + name = call.get("name") + if not isinstance(name, str) or not name.strip(): + continue + + raw_args = call.get("arguments", call.get("args", {})) + if isinstance(raw_args, str): + args_json = raw_args + else: + safe_args = raw_args if isinstance(raw_args, dict) else {} + args_json = json.dumps(safe_args, ensure_ascii=False) + + normalized_call: dict[str, Any] = { + "type": "function", + "function": { + "name": name, + "arguments": args_json, + }, + } + call_id = call.get("id") or call.get("tool_call_id") + if isinstance(call_id, str) and call_id.strip(): + normalized_call["id"] = call_id + normalized.append(normalized_call) + return normalized + + +def _extract_tool_calls(message: Any) -> list[dict[str, Any]]: + tool_calls = getattr(message, "tool_calls", None) + if not tool_calls: + return [] + + normalized: list[dict[str, Any]] = [] + for call in tool_calls: + try: + name = call.function.name + arguments_raw = call.function.arguments + try: + arguments = json.loads(arguments_raw) if arguments_raw else {} + except json.JSONDecodeError: + arguments = {"raw": arguments_raw} + normalized.append( + { + "id": getattr(call, "id", None), + "name": name, + "arguments": arguments, + } + ) + except AttributeError: + continue + return normalized + + +def _build_async_http_client(options: dict[str, Any]) -> Any | None: + if not options: + return None + try: + import httpx + except Exception: + return None + try: + return httpx.AsyncClient(**options) + except Exception: + return None + + +def _extract_thinking_content(message: Any) -> str | None: + for attr in ("reasoning_content", "reasoning", "thinking"): + text = _coerce_text(getattr(message, attr, None)) + if text: + return text + + additional_kwargs = getattr(message, "additional_kwargs", None) + if isinstance(additional_kwargs, dict): + for key in ("reasoning_content", "reasoning", "thinking"): + text = _coerce_text(additional_kwargs.get(key)) + if text: + return text + + model_extra = getattr(message, "model_extra", None) + if isinstance(model_extra, dict): + for key in ("reasoning_content", "reasoning", "thinking"): + text = _coerce_text(model_extra.get(key)) + if text: + return text + return None + + +def _extract_reasoning_tokens(response: Any) -> int | None: + usage = getattr(response, "usage", None) + if usage is None: + return None + candidates: list[Any] = [ + getattr(usage, "reasoning_tokens", None), + _get_nested_value(getattr(usage, "completion_tokens_details", None), "reasoning_tokens"), + _get_nested_value(getattr(usage, "output_tokens_details", None), "reasoning_tokens"), + _get_nested_value(getattr(usage, "output_tokens_details", None), "reasoning"), + ] + for candidate in candidates: + if candidate is None: + continue + try: + return int(candidate) + except (TypeError, ValueError): + continue + return None + + +def _get_nested_value(value: Any, key: str) -> Any: + if isinstance(value, dict): + return value.get(key) + return getattr(value, key, None) + + +def _coerce_text(value: Any) -> str | None: + if isinstance(value, str): + text = value.strip() + return text or None + if isinstance(value, dict): + for key in ("text", "content", "reasoning", "thinking"): + text = _coerce_text(value.get(key)) + if text: + return text + return None + if isinstance(value, list): + parts: list[str] = [] + for item in value: + text = _coerce_text(item) + if text: + parts.append(text) + if parts: + return "\n".join(parts) + return None + + +__all__ = ["HuaweiModelArtsModelAdapter"] diff --git a/dare_framework/model/default_model_adapter_manager.py b/dare_framework/model/default_model_adapter_manager.py index 71f2b09b..4cb09217 100644 --- a/dare_framework/model/default_model_adapter_manager.py +++ b/dare_framework/model/default_model_adapter_manager.py @@ -8,6 +8,7 @@ from dare_framework.model.interfaces import IModelAdapterManager from dare_framework.model.kernel import IModelAdapter from dare_framework.model.adapters.anthropic_adapter import AnthropicModelAdapter +from dare_framework.model.adapters.huawei_modelarts_adapter import HuaweiModelArtsModelAdapter from dare_framework.model.adapters.openai_adapter import OpenAIModelAdapter from dare_framework.model.adapters.openrouter_adapter import OpenRouterModelAdapter @@ -30,8 +31,10 @@ def load_model_adapter(self, *, config: Config | None = None) -> IModelAdapter | return _build_openrouter_adapter(llm) if adapter_name == "anthropic": return _build_anthropic_adapter(llm) + if adapter_name == "huawei-modelarts": + return _build_huawei_modelarts_adapter(llm) raise ValueError( - f"Unsupported model adapter '{adapter_name}'. Supported adapters: openai, openrouter, anthropic." + f"Unsupported model adapter '{adapter_name}'. Supported adapters: openai, openrouter, anthropic, huawei-modelarts." ) @@ -75,6 +78,17 @@ def _build_anthropic_adapter(llm: LLMConfig) -> AnthropicModelAdapter: ) +def _build_huawei_modelarts_adapter(llm: LLMConfig) -> HuaweiModelArtsModelAdapter: + return HuaweiModelArtsModelAdapter( + name="huawei-modelarts", + api_key=llm.api_key, + model=llm.model, + base_url=llm.endpoint, + http_client_options=_http_client_options_from_proxy(llm), + extra=dict(llm.extra), + ) + + def _http_client_options_from_proxy(llm: LLMConfig) -> dict[str, Any]: proxy = llm.proxy options: dict[str, Any] = {} diff --git a/tests/unit/test_client_cli.py b/tests/unit/test_client_cli.py index cba28172..0964ac53 100644 --- a/tests/unit/test_client_cli.py +++ b/tests/unit/test_client_cli.py @@ -351,6 +351,25 @@ def test_build_doctor_report_accepts_anthropic_with_key() -> None: assert not any("unsupported adapter configured" in item for item in payload["warnings"]) +def test_build_doctor_report_accepts_huawei_modelarts_with_key() -> None: + config = Config.from_dict( + { + "workspace_dir": ".", + "user_dir": ".", + "llm": { + "adapter": "huawei-modelarts", + "model": "modelarts-pro", + "api_key": "dummy", + }, + "mcp_paths": [], + } + ) + payload = build_doctor_report(config=config) + assert payload["llm"]["adapter"] == "huawei-modelarts" + assert payload["llm"]["api_key_present"] is True + assert not any("unsupported adapter configured" in item for item in payload["warnings"]) + + @pytest.mark.asyncio async def test_main_doctor_does_not_bootstrap_runtime(monkeypatch, tmp_path) -> None: client_main = importlib.import_module("client.main") diff --git a/tests/unit/test_default_model_adapter_manager.py b/tests/unit/test_default_model_adapter_manager.py index 048db1fc..f796b520 100644 --- a/tests/unit/test_default_model_adapter_manager.py +++ b/tests/unit/test_default_model_adapter_manager.py @@ -4,7 +4,12 @@ from dare_framework.agent import BaseAgent from dare_framework.config.types import Config, LLMConfig -from dare_framework.model import AnthropicModelAdapter, OpenAIModelAdapter, OpenRouterModelAdapter +from dare_framework.model import ( + AnthropicModelAdapter, + HuaweiModelArtsModelAdapter, + OpenAIModelAdapter, + OpenRouterModelAdapter, +) from dare_framework.model.default_model_adapter_manager import DefaultModelAdapterManager @@ -33,6 +38,21 @@ def test_default_manager_returns_anthropic_adapter() -> None: assert adapter.model_name == "claude-sonnet-4-5" +def test_default_manager_returns_huawei_modelarts_adapter() -> None: + manager = DefaultModelAdapterManager() + config = Config( + llm=LLMConfig( + adapter="huawei-modelarts", + api_key="test-key", + model="modelarts-pro", + ) + ) + adapter = manager.load_model_adapter(config=config) + assert isinstance(adapter, HuaweiModelArtsModelAdapter) + assert adapter.name == "huawei-modelarts" + assert adapter.model_name == "modelarts-pro" + + def test_default_manager_unsupported_adapter_raises() -> None: manager = DefaultModelAdapterManager() config = Config(llm=LLMConfig(adapter="unknown")) diff --git a/tests/unit/test_huawei_modelarts_adapter.py b/tests/unit/test_huawei_modelarts_adapter.py new file mode 100644 index 00000000..47cf80eb --- /dev/null +++ b/tests/unit/test_huawei_modelarts_adapter.py @@ -0,0 +1,143 @@ +from __future__ import annotations + +from types import SimpleNamespace + +import pytest + +from dare_framework.context.types import AttachmentKind, AttachmentRef, Message +from dare_framework.model.adapters.huawei_modelarts_adapter import ( + HuaweiModelArtsModelAdapter, + _extract_reasoning_tokens, + _extract_thinking_content, + _serialize_messages, +) + + +def test_adapter_requires_api_key() -> None: + with pytest.raises(ValueError, match="API key is required"): + HuaweiModelArtsModelAdapter(api_key=None, model="modelarts-pro") + + +def test_adapter_requires_model() -> None: + with pytest.raises(ValueError, match="model is required"): + HuaweiModelArtsModelAdapter(api_key="dummy", model=None) + + +def test_serialize_messages_preserves_assistant_tool_calls() -> None: + messages = [ + Message(role="user", text="Need a filename"), + Message( + role="assistant", + text="I need your confirmation first.", + data={ + "tool_calls": [ + { + "id": "call_1", + "name": "ask_user", + "arguments": { + "questions": [ + { + "header": "Target", + "question": "Pick a file name", + "options": [ + {"label": "a.txt", "description": "A"}, + {"label": "b.txt", "description": "B"}, + ], + } + ] + }, + } + ] + }, + ), + Message( + role="tool", + name="call_1", + text='{"success": true, "output": {"answers": {"Pick a file name": "a.txt"}}}', + ), + ] + + serialized = _serialize_messages(messages) + + assistant_payload = serialized[1] + assert "tool_calls" in assistant_payload + tool_call = assistant_payload["tool_calls"][0] + assert tool_call["id"] == "call_1" + assert tool_call["type"] == "function" + assert tool_call["function"]["name"] == "ask_user" + assert isinstance(tool_call["function"]["arguments"], str) + assert serialized[2]["tool_call_id"] == "call_1" + + +def test_serialize_messages_supports_chat_text_with_image_attachments() -> None: + messages = [ + Message( + role="user", + text="describe both", + attachments=[ + AttachmentRef(kind=AttachmentKind.IMAGE, uri="https://example.com/a.png"), + AttachmentRef(kind=AttachmentKind.IMAGE, uri="https://example.com/b.png"), + ], + ) + ] + + serialized = _serialize_messages(messages) + + assert serialized == [ + { + "role": "user", + "content": [ + {"type": "text", "text": "describe both"}, + {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}}, + {"type": "image_url", "image_url": {"url": "https://example.com/b.png"}}, + ], + } + ] + + +def test_serialize_messages_prefers_structured_data_for_tool_history() -> None: + messages = [ + Message( + role="assistant", + text="call tool", + kind="tool_call", + data={ + "tool_calls": [ + { + "id": "call_1", + "name": "search", + "arguments": {"q": "docs"}, + } + ] + }, + ), + Message( + role="tool", + kind="tool_result", + text="done", + data={"tool_call_id": "call_1"}, + ), + ] + + serialized = _serialize_messages(messages) + + assert serialized[0]["tool_calls"][0]["id"] == "call_1" + assert serialized[1]["tool_call_id"] == "call_1" + + +def test_extract_thinking_content_from_message_fields() -> None: + message = SimpleNamespace(reasoning="step by step", reasoning_content=None, additional_kwargs={}) + assert _extract_thinking_content(message) == "step by step" + + +def test_extract_reasoning_tokens_from_completion_tokens_details() -> None: + response = SimpleNamespace( + usage=SimpleNamespace( + prompt_tokens=1, + completion_tokens=2, + total_tokens=3, + reasoning_tokens=None, + completion_tokens_details=SimpleNamespace(reasoning_tokens=7), + ) + ) + assert _extract_reasoning_tokens(response) == 7