From 65fcf2b82f87f88dd299cf25de919247b4f903fa Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 10:20:34 +0800 Subject: [PATCH 01/55] =?UTF-8?q?docs:=20AIBridge=20Rust=20=E9=87=8D?= =?UTF-8?q?=E6=9E=84=E8=AE=BE=E8=AE=A1=E6=96=87=E6=A1=A3?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 将 agn-sdk 用 Rust 重构为跨语言 SDK aibridge 的设计文档。 - 五语言(Python/JS/Go/JVM/.NET)原生 import - Rust 核心 + 双入口绑定(PyO3/napi 直连 + C ABI) - v2 破坏性升级,MVP 优先 agnes/火山/gemini/openai --- ...2026-07-07-aibridge-rust-rewrite-design.md | 459 ++++++++++++++++++ 1 file changed, 459 insertions(+) create mode 100644 docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md diff --git a/docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md b/docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md new file mode 100644 index 0000000..60130e7 --- /dev/null +++ b/docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md @@ -0,0 +1,459 @@ +# AIBridge Rust 重构设计文档 + +> 日期:2026-07-07 +> 状态:草案,待用户审查 +> 范围:将现有 Python `agn-sdk` (v1.3.3, ~19700 行) 用 Rust 重构为跨语言 SDK `aibridge` + +--- + +## 1. 背景与目标 + +AIBridge(原名 agn-sdk,因早期基于 Agnes 开发而得名)是多模态 AI 统一接口 SDK,一套 API 调用所有 AI 模型(chat / image / video / TTS / ASR / embed)。现有 Python 实现已发布 PyPI,已有项目在用。 + +本次重构的核心目标: + +1. **五语言原生 import**:Python / JS-TS / Go / JVM(Java,Kotlin) / .NET(C#) 直接调用同一套能力 +2. **全量迁移**:14 个适配器 + 6 大能力全部用 Rust 重写 +3. **MVP 优先**:agnes / 火山引擎(volcengine_cv) / gemini / openai 四个 provider 先在五语言跑通(用户另一项目主力依赖) +4. **v2 破坏性升级**:重新设计 FFI 友好的统一 API,提供 Python v1→v2 迁移指南 +5. **更名与归档**:`agn-sdk` → `aibridge`,旧版 v1 归档,新版接管 + +非目标: + +- 不保留 v1 API 1:1 兼容(用迁移指南过渡) +- 不做 Python 老版性能基准对照 +- 不在首版引入新 provider(仅迁移现有 14 个) + +--- + +## 2. 关键决策摘要 + +| # | 决策点 | 选择 | 理由 | +|---|---|---|---| +| 1 | 目标语言 | Python + JS/TS + Go + JVM + .NET | 用户要求五语言全覆盖 | +| 2 | Python 兼容性 | v2 破坏性升级 + 迁移指南 | 1:1 兼容会绑架所有语言 API 风格,成本过高 | +| 3 | 首版范围 | 全量迁移,MVP 优先四 provider | 全量是最终目标,四 provider 先跑通验证管线 | +| 4 | 流式/异步 | 全原生 async | 体验最佳;PyO3/napi 直连 + C ABI 静态语言包装 | +| 5 | 仓库/发布 | Monorepo + 接管更名 | 统一版本/CI,semver v2.0.0 表达破坏性 | +| 6 | 架构方案 | Rust 核心 + 双入口绑定 | 唯一同时满足全原生 async + 五语言体验一致 | +| 7 | 品牌名 | aibridge | "AI 桥",直白好记,PyPI/npm 均可用 | + +--- + +## 3. 整体架构 + +``` + aibridge-core (Rust, 纯 async 逻辑) + ┌──────────┴──────────┐ + 直连(原生async) C ABI (aibridge-ffi cdylib) + ┌─────┴─────┐ ┌─────┬─────┬─────┐ + aibridge- aibridge- aibridge- aibridge- aibridge- + python node go jvm dotnet + (PyO3) (napi-rs) (CGO) (JNA) (P/Invoke) + asyncio Promise/ goroutine CompletableFuture Task/ + AsyncIter AsyncIter +channel /Flow IAsyncEnum +``` + +**双入口原理**:Python/JS 用 PyO3/napi-rs 直连 `aibridge-core`,享真正原生 async(无序列化边界);Go/JVM/.NET 没有跨语言 async FFI 概念,走 `aibridge-ffi` 的 C ABI(阻塞调用 + 各语言原生异步原语包装)。五种语言共享同一个 Rust 核心,绑定层都薄。 + +**保持原五层心智**:API 层(Client) → 路由层(Router) → 适配器层(Adapter) → 核心层(Core) → 模型层(Model),仅语言换 Rust、`**kwargs` 换显式 struct、Pydantic 换 serde、动态注册换编译期 match。 + +--- + +## 4. Monorepo 布局 + +``` +aibridge/ # 重构现有 agn-sdk 仓库并更名 +├── Cargo.toml # workspace 根 +├── crates/ +│ ├── aibridge-core/ # Rust 核心(纯 async 逻辑,无 FFI 污染) +│ │ ├── src/ +│ │ │ ├── lib.rs +│ │ │ ├── client.rs # 统一 Client +│ │ │ ├── router.rs # Router 路由/fallback +│ │ │ ├── adapter/{mod,base,factory}.rs # Adapter trait + 工厂 +│ │ │ ├── adapters/ # 14 个适配器 +│ │ │ │ ├── openai_compat.rs # OpenAI 兼容地基 +│ │ │ │ ├── openai.rs agnes.rs azure.rs +│ │ │ │ ├── anthropic.rs gemini.rs volcengine_cv.rs +│ │ │ │ ├── runway.rs pika.rs kling.rs stability.rs +│ │ │ │ ├── chinese.rs aggregation_platforms.rs +│ │ │ │ ├── additional_models.rs more_models.rs emerging_models.rs +│ │ │ │ └── audio_adapters.rs # edge-tts/elevenlabs/cartesia/deepgram/assemblyai +│ │ │ ├── http.rs # reqwest+h2 封装 +│ │ │ ├── retry.rs # 重试机制 +│ │ │ ├── error.rs # thiserror 错误枚举 +│ │ │ ├── config.rs # 配置 + 环境变量 +│ │ │ ├── util.rs +│ │ │ └── model/{chat,image,video,audio,common,options}.rs # serde struct +│ │ └── tests/ +│ ├── aibridge-ffi/ # C ABI cdylib(给 Go/JVM/.NET) +│ │ ├── src/{lib,handle,error,stream,runtime}.rs +│ │ ├── include/aibridge.h # cbindgen 生成 +│ │ └── cbindgen.toml +│ ├── aibridge-python/ # PyO3 绑定(直连 aibridge-core) +│ │ ├── src/lib.rs # #[pymodule] aibridge +│ │ └── pyproject.toml # maturin +│ └── aibridge-node/ # napi-rs 绑定(直连 aibridge-core) +│ ├── src/lib.rs +│ ├── package.json +│ └── build.rs +├── bindings/ +│ ├── go/ # CGO 调 aibridge-ffi(goroutine+channel) +│ │ ├── aibridge.go aibridge.go.h stream.go +│ │ └── go.mod (github.com/aibridge/aibridge-go) +│ ├── jvm/ # JNA 调 aibridge-ffi(CompletableFuture/Flow,纯 Java 无 native Rust) +│ │ ├── src/main/{java,kotlin}/io/aibridge/... +│ │ └── build.gradle.kts +│ └── dotnet/ # C# P/Invoke 调 aibridge-ffi(Task/IAsyncEnumerable) +│ ├── AIBridge/ (C# 项目) +│ └── AIBridge.csproj +├── tests/ # 跨语言一致性测试套件 + 共享 fixture +├── docs/ # 设计文档 + 各语言迁移指南 +└── .github/workflows/ # CI 矩阵(5 语言 × 多平台) +``` + +**关键决策**:Go/JVM/.NET 统一走 `aibridge-ffi` 这一个 C ABI 入口。JVM 用 JNA(纯 Java,无需写 native Rust),Go 用 CGO,.NET 用 P/Invoke。避免三种静态语言各写一套 Rust native 胶水。 + +--- + +## 5. aibridge-core 模块设计 + +### 5.1 Python → Rust 模块映射 + +| Python 现有 | Rust (aibridge-core) | 说明 | +|---|---|---| +| `agn/client.py` | `client::Client` | 统一入口,方法签名改显式 Request struct | +| `agn/router.py` | `router::Router` | 多 provider 路由/负载均衡/fallback | +| `adapters/base.py` BaseAdapter | `adapter::Adapter` trait | `#[async_trait]`,不支持能力默认返 UnsupportedCapability | +| `adapters/factory.py` AdapterFactory | `adapter::create_adapter()` 显式 match | 替代运行时注册,新增适配器=加 match 分支 | +| `adapters/*.py`(14个) | `adapters::*`(14个模块) | trait 实现 | +| `core/http_client.py` | `http` (reqwest+h2) | 替代 httpx | +| `core/retry.py` | `retry` | 替代 tenacity | +| `core/errors.py` | `error` (thiserror) | 标准错误枚举 | +| `core/config.py` | `config` | 配置 + 环境变量 | +| `models/*.py` (Pydantic) | `model::*` (serde struct) | 替代 Pydantic | + +### 5.2 Adapter trait + +```rust +#[async_trait] +pub trait Adapter: Send + Sync { + fn provider_type(&self) -> &str; + fn provider_name(&self) -> &str; + fn capabilities(&self) -> Capabilities; + fn requires_api_key(&self) -> bool { true } + + async fn start(&mut self) -> Result<()>; + async fn close(&mut self) -> Result<()>; + async fn chat(&self, req: ChatRequest) -> Result; + async fn chat_stream(&self, req: ChatRequest) -> Result; // impl Stream + async fn image_generate(&self, req: ImageRequest) -> Result; + async fn video_create(&self, req: VideoRequest) -> Result; + async fn video_poll(&self, task_id: &str, model: &str) -> Result; + async fn embed(&self, req: EmbedRequest) -> Result; + async fn transcribe(&self, req: TranscribeRequest) -> Result; + async fn speech(&self, req: SpeechRequest) -> Result; + async fn list_models(&self, filter: Option) -> Result>; + async fn list_voices(&self, language: Option<&str>) -> Result>; + async fn recommend_voices(&self, language: Option<&str>, gender: Option<&str>, limit: usize) -> Result>; + + // 不支持的方法默认返 UnsupportedCapabilityError,trait 提供默认实现 +} +``` + +### 5.3 Client 形态 + +```rust +pub struct Client { adapter: Box } +impl Client { + pub fn new(provider: &str, opts: ClientOptions) -> Result; + pub async fn chat(&self, req: ChatRequest) -> Result; + pub async fn chat_stream(&self, req: ChatRequest) -> Result; + pub async fn image_generate(&self, req: ImageRequest) -> Result; + pub async fn video_create(&self, req: VideoRequest) -> Result; + pub async fn video_poll(&self, task_id: &str, model: &str) -> Result; + pub async fn embed(&self, req: EmbedRequest) -> Result; + pub async fn transcribe(&self, req: TranscribeRequest) -> Result; + pub async fn speech(&self, req: SpeechRequest) -> Result; + pub async fn list_models(&self, filter: Option) -> Result>; + pub async fn list_voices(&self, language: Option<&str>) -> Result>; + pub async fn recommend_voices(&self, lang: Option<&str>, gender: Option<&str>, limit: usize) -> Result>; +} +``` + +- **可选参数**:Request struct + Builder(`..Default::default()`),替代 `**kwargs`。Provider 特有参数走 `extra: HashMap` 透传。 +- **流式**:`ChatStream: impl Stream>`(`async-stream` crate),核心内部原生 async stream。 +- **工厂**:编译期显式 `match`,替代 Python 运行时 `AdapterFactory.register`。更静态、更安全。 + +--- + +## 6. 统一数据模型(serde struct 替代 Pydantic + kwargs) + +```rust +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct ChatRequest { + pub model: String, + pub messages: Vec, + pub temperature: Option, + pub max_tokens: Option, + pub stop: Option, + pub tools: Option>, + pub tool_choice: Option, + pub reasoning_effort: Option, + pub response_format: Option, + pub extra: HashMap, // provider 特有参数透传 +} +// ChatRequest::builder("gpt-4o", messages).temperature(0.7).max_tokens(1000).build() + +#[serde(tag = "role", rename_all = "lowercase")] +pub enum ChatMessage { + System { content: String }, + User { content: UserContent }, // String 或多模态 Vec + Assistant { content: String, tool_calls: Option> }, + Tool { tool_call_id: String, content: String }, +} + +pub enum FileInput { Path(String), Url(String), Bytes(Vec), Base64(String) } +``` + +- **去掉 Options 中间层**:Python 的 `ChatOptions/ImageOptions/...` 在 Rust 直接用 `Request::builder()` 链式调用替代,迁移指南做 1:1 对照。 +- **响应模型**:`ChatCompletion`/`ChatCompletionChunk`/`ImageResult`/`VideoTask`/`VideoStatus`/`EmbeddingResult`/`TranscriptionResult`/`SpeechResult`/`ModelInfo`/`VoiceInfo` 全部 serde struct,字段名与 Python Pydantic 对齐(迁移指南可逐字段对照)。 + +--- + +## 7. aibridge-ffi C ABI 边界 + +**核心原则**:句柄式生命周期 + JSON 字符串作复杂 struct 边界 + 二进制原生传递 + 全局 tokio runtime。 + +```c +/* 句柄(opaque) */ +typedef struct aibridge_client_s aibridge_client_t; +typedef struct aibridge_stream_s aibridge_stream_t; +typedef struct { const uint8_t* ptr; size_t len; } aibridge_bytes_t; + +/* 生命周期 */ +aibridge_client_t* aibridge_client_new(const char* provider, const char* config_json); +aibridge_status_t aibridge_client_start(aibridge_client_t*); +void aibridge_client_destroy(aibridge_client_t*); + +/* 阻塞式调用:内部 runtime.block_on,复杂结构走 JSON 边界 */ +aibridge_status_t aibridge_client_chat(aibridge_client_t*, const char* request_json, + char** out_response_json); /* 调用方 aibridge_string_free */ +/* 二进制载荷不走 JSON(避免 base64 膨胀) */ +aibridge_status_t aibridge_client_speech(aibridge_client_t*, const char* request_json, + aibridge_bytes_t** out_audio, char** out_meta_json); + +/* 流式:stream 句柄 + 阻塞 next() */ +aibridge_status_t aibridge_client_chat_stream(aibridge_client_t*, const char* request_json, + aibridge_stream_t** out_stream); +aibridge_status_t aibridge_stream_next(aibridge_stream_t*, char** out_chunk_json); /* 0=chunk, 1=结束, 负=错 */ +void aibridge_stream_destroy(aibridge_stream_t*); /* 触发 Rust drop → tokio task abort */ + +/* 错误:返回码 + 线程局部 last_error 槽 */ +const char* aibridge_last_error(void); /* 线程局部,无需释放 */ + +/* 释放 */ +void aibridge_string_free(char*); +void aibridge_bytes_free(aibridge_bytes_t*); +``` + +**关键决策与权衡**: + +1. **JSON 边界 vs C struct mirrorgen**:选 JSON。5 语言 × 大量 struct 用 mirrorgen 维护成本高;JSON 边界让 C ABI 只需 ~20 个函数,各语言绑定只做 JSON (de)serialize,类型安全且跨语言一致。SDK 是 IO 密集,JSON 序列化开销可忽略。 +2. **全局 tokio runtime**:`Lazy`(多线程),每个 FFI 调用 `handle().block_on(async{...})`。C 侧可并发调用不同 client;同一 stream 串行(各语言绑定负责同步)。 +3. **错误模型**:`aibridge_status_t` 返回码(0 成功 / 负数错误类别)+ `aibridge_last_error()` 线程局部详细消息(JSON,含 `code/message/details/retryable`)。各语言绑定映射成本地异常。 +4. **字符串/二进制所有权**:Rust 分配的 `char*`/`aibridge_bytes_t` 必须由调用方调 `aibridge_*_free` 释放。各语言绑定用 RAII 封装(Go finalizer / JVM Cleaner / .NET Dispose)。 + +--- + +## 8. 异步与流式跨语言桥接 + +| 语言 | 绑定路径 | 异步原语 | 流式类型 | cancel 机制 | +|---|---|---|---|---| +| Python | PyO3 **直连** aibridge-core | asyncio 协程 | `AsyncIterator` | asyncio cancel → drop stream | +| JS/TS | napi-rs **直连** aibridge-core | `Promise` | `AsyncIterable` | `AbortSignal` → drop stream | +| Go | CGO → aibridge-ffi | goroutine | `chan Chunk` | `ctx.Done()` → `aibridge_stream_destroy` | +| JVM | JNA → aibridge-ffi | `CompletableFuture` / Kotlin `suspend` | `Flow.Publisher` | cancel → destroy | +| .NET | P/Invoke → aibridge-ffi | `Task` | `IAsyncEnumerable` | `CancellationToken` → destroy | + +**直连 vs C ABI**: + +- **Python/JS 直连**:无 JSON 边界,PyO3 `#[pyo3] async fn` / napi `#[napi] async fn` 自动桥接原生 async;Rust struct 用 `#[pyclass]`/napi 标注后直接映射为宿主对象。流式内部 `tokio::spawn` + channel 桥接成原生异步迭代器。体验最佳、零序列化。 +- **Go/JVM/.NET 走 C ABI**:FFI 调用本身阻塞,各语言在自己的异步执行器上调度(Go goroutine / JVM ForkJoinPool / .NET ThreadPool),把阻塞 FFI 包成原生异步原语。流式用"stream 句柄 + 阻塞 `next()`"在后台线程循环,结果 push 到 channel/Flow/IAsyncEnumerable。 +- **统一 cancel**:所有语言的取消最终都落到 `aibridge_stream_destroy`(Rust drop → tokio task abort),语义一致。Python/JS 的直连绑定把 asyncio cancel / AbortSignal 桥接为 Rust 侧 drop stream。 + +--- + +## 9. 错误处理与类型映射 + +### 9.1 Rust 错误枚举(thiserror) + +```rust +#[derive(Debug, thiserror::Error)] +pub enum AibridgeError { + #[error("认证失败: {message}")] Authentication { message: String }, + #[error("限流: {message}")] RateLimit { message: String, retry_after: Option }, + #[error("参数校验错误: {message}")] Validation { message: String, details: serde_json::Value }, + #[error("模型不存在: {model}")] ModelNotFound { model: String }, + #[error("API 调用错误: {message}")] Api { status: u16, message: String }, + #[error("网络错误: {0}")] Network(#[from] reqwest::Error), + #[error("超时")] Timeout, + #[error("不支持的能力: {capability}")] UnsupportedCapability { capability: String }, + #[error("Provider 不存在: {provider}")] ProviderNotFound { provider: String }, +} +``` + +对应 Python v1 的 `AGNError` 体系(`AuthenticationError`/`RateLimitError`/`ValidationError`/`ModelNotFoundError`/`APIError`/`TimeoutError`/`UnsupportedCapabilityError`/`ProviderNotFoundError`),分类一致,迁移指南做名称对照(`AGNError` → `AibridgeError`)。 + +### 9.2 FFI 错误传递 + +返回码 `aibridge_status_t`(i32)映射错误类别;`aibridge_last_error()` 返回线程局部 JSON:`{"code":"rate_limit","message":"...","retryable":true,"details":{...}}`。 + +### 9.3 各语言异常映射 + +- **Python**:`AibridgeError` 基类 + 子类(`RateLimitError` 等),与 v1 同名便于迁移 +- **JS/TS**:`AibridgeError` + 子类 +- **Go**:`error` + 类型断言接口 `type AibridgeError interface{ Code() string }` +- **JVM**:`AibridgeException` + 子类 +- **.NET**:`AibridgeException` + 子类 + +--- + +## 10. 适配器迁移策略 + +### 10.1 抓共性:OpenAI 兼容地基 + +14 个适配器中 80% 是 OpenAI 兼容协议(agnes/openai/azure/聚合平台/中文平台兼容部分)。Rust 抽出 `OpenAiCompatAdapter` 基础实现(HTTP 请求构造 + 响应解析 + 参数映射),子适配器只 override 差异(base_url、model 列表、参数 mapping、特殊端点)。 + +### 10.2 参数映射机制 Rust 化 + +保留 Python `ParameterMapping` 概念,Rust 化为 `param_mapping` 表 + `apply_mapping()`。预置常量:`OPENAI_COMPATIBLE_MAPPING`、`ANTHROPIC_MAPPING`、`GEMINI_MAPPING`、`COHERE_MAPPING`,OpenAI 兼容适配器复用。 + +### 10.3 分批迁移顺序(MVP 优先) + +| 批次 | 适配器 | 协议 | 优先级 | +|---|---|---|---| +| **阶段 1 · MVP** | openai, agnes, volcengine_cv, gemini | 兼容地基 + 独立协议 | **最高** | +| 阶段 2a | azure, aggregation_platforms, additional_models, more_models, emerging_models, chinese(兼容部分) | OpenAI 兼容族 | 高 | +| 阶段 2b | anthropic, stability, runway, pika, kling | 独立协议 | 中 | +| 阶段 2c | edge-tts, elevenlabs, cartesia, deepgram, assemblyai (audio_adapters) | 二进制载荷 | 中 | + +> 工作量估算:~14000 行 Python,共享地基后实际新写 Rust 约 6000–8000 行。 + +### 10.4 保留特性 + +- v1.1.0 的"实时拉取模型列表"(list_models 调 provider /models 端点)保留 +- v1.3.0 的"免费 Provider 免认证"(edge-tts 无需 api_key,`requires_api_key() -> false`)保留 +- v1.3.3 的"TTS 音色健康检查/推荐/自动降级"保留 + +--- + +## 11. 测试策略 + +三层测试: + +1. **Rust 核心单测**(aibridge-core 内):每个适配器 mock HTTP,覆盖正常 + 异常 + 边界。覆盖率 ≥80%。 +2. **各语言绑定测试**:每种语言测自己的胶水层(句柄生命周期、异步桥接、错误映射、流式)。 +3. **跨语言一致性测试**:同一组输入 + mock,五种语言跑一遍,断言输出一致。防止单种语言绑定漏实现或行为漂移。 + +**共享 fixture**:Python v1 测试里的 HTTP 请求/响应 mock 数据抽成 JSON 文件(`tests/fixtures/`),Rust 测试 + 五语言绑定测试共用。保证"五语言行为一致 + 与 Python 老版一致"。**这是质量保险的核心。** + +平时测试不调真实 AI API(成本高、不稳定)。另留"真实接口冒烟测试"开关(环境变量配 key,CI 默认关)。 + +--- + +## 12. 构建与发布 + +### 12.1 各语言打包方式 + +| 语言 | 工具 | 发布到 | 用户安装体验 | +|---|---|---|---| +| Rust 核心 + ffi | cargo | —(内部) | 产出 libaibridge.{so,dylib,dll} | +| Python | maturin | PyPI `aibridge` v2.0.0 | `pip install aibridge`,预编译 wheel,无需 Rust | +| JS/TS | napi-rs | npm `aibridge` | `npm install aibridge`,预编译 .node,无需 Rust | +| Go | cgo | Go module `aibridge-go` | 需单独装 libaibridge(提供安装脚本,Go 生态惯例) | +| JVM | JNA + Gradle | Maven `io.aibridge:aibridge` | 动态库打进 jar(按平台 classifier),用户无感 | +| .NET | P/Invoke | NuGet `AIBridge` | 动态库打进包(runtimes/{rid}/native/),用户无感 | + +### 12.2 二进制分发难点 + +Go/JVM/.NET 依赖 `libaibridge` 动态库。JVM 和 .NET 把动态库打进各自包里(按 OS/arch 分包),用户无感;Go 因 cgo 机制,用户单独装动态库(提供安装脚本),属 Go 生态常规做法。 + +### 12.3 CI 矩阵 + +GitHub Actions:平台(linux/macos/windows × amd64/arm64)× 语言。aibridge-ffi 的动态库作为构建 artifact 供 Go/JVM/.NET 打包消费。Rust 核心 + 5 绑定各一个 workflow。 + +--- + +## 13. 分阶段实施计划 + +> 时间为全职单人粗估;实施阶段用多 agent 并行可加速。 + +### 阶段 0 · 打地基(2–3 周) + +- Monorepo 搭建,Cargo workspace +- aibridge-core 骨架:error/config/http/retry/model/Adapter trait/Client/Router +- aibridge-ffi C ABI 骨架 + 全局 runtime + cbindgen +- 跨语言管线打通:用 openai 的 chat stub 跑通五语言 hello world(含 async + 流式 + 错误),验证最大技术风险 + +### 阶段 1 · MVP 四 provider(3–4 周)⭐ + +openai → agnes → volcengine_cv → gemini,逐个 Rust 实现并跑通五语言 + 测试。**做完这步,用户另一项目即可开始接入。** + +### 阶段 2 · 剩余适配器(4–6 周) + +按 2a(兼容族) → 2b(独立协议) → 2c(音频) 三批搬完剩下 10 个。 + +### 阶段 3 · 发布收尾(2–3 周) + +全平台 CI 构建、五语言包正式发版、Python v1→v2 迁移指南、文档网站、老版 v1 归档、打 v2.0.0 tag。 + +**总周期约 3–4 个月**(全职单人)。 + +### 多 agent 并行编排策略(实施阶段) + +- **适配器迁移**:每个 provider 一个 agent 并行(共享 OpenAiCompatAdapter 基础),fixture 共享保证一致 +- **语言绑定**:aibridge-core 与 aibridge-ffi 稳定后,五种语言绑定各一个 agent 并行 +- **测试**:跨语言一致性测试单独 agent 汇总 +- **依赖顺序**:阶段 0 必须串行(核心未稳定前绑定无法并行);阶段 1+ 适配器与绑定可并行 + +--- + +## 14. 风险与缓解 + +| 风险 | 影响 | 缓解 | +|---|---|---| +| 全原生 async 跨 FFI 复杂 | 高 | PyO3/napi 直连避开 C ABI async;静态语言用阻塞+原生异步包装,已在设计中固化 | +| 14 适配器全量迁移周期长 | 中 | OpenAI 兼容地基复用 80%;MVP 四 provider 优先保证早期可用 | +| 五语言行为漂移 | 中 | 共享 fixture + 跨语言一致性测试 | +| Go/JVM/.NET 动态库分发 | 中 | JVM/.NET 打进包;Go 提供安装脚本 | +| crates.io/包名占用 | 低 | PyPI/npm 已确认 aibridge 可用;crates.io/Maven/NuGet 发布前再确认,可用后缀规避 | +| Python 老用户升级破坏 | 中 | v2.0.0 semver + 迁移指南 + 旧版 v1 归档保留 | + +--- + +## 15. Python v1→v2 迁移指南要点 + +- 包名:`agn-sdk` → `aibridge`,`from agn import Client` → `from aibridge import Client` +- 错误类:`AGNError` → `AibridgeError`(子类名不变) +- 参数:`**kwargs` 透传 → `Request` struct + Builder 链式调用(`ChatOptions` → `ChatRequest::builder()`) +- Options 类:`ChatOptions/ImageOptions/...` 中间层去除,直接用 Request builder +- 其余方法名、能力、provider 名保持一致 + +--- + +## 附录 A:与 Python v1 能力对照 + +| 能力 | v1 (agn-sdk) | v2 (aibridge) | 状态 | +|---|---|---|---| +| chat (含流式) | ✅ | ✅ | 迁移 | +| image_generate | ✅ | ✅ | 迁移 | +| video_create + poll | ✅ | ✅ | 迁移 | +| transcribe (ASR) | ✅ | ✅ | 迁移 | +| speech (TTS) | ✅ | ✅ | 迁移 | +| embed | ✅ | ✅ | 迁移 | +| list_models (实时拉取) | ✅ | ✅ | 迁移 | +| list_voices / recommend_voices | ✅ | ✅ | 迁移 | +| Router (多 provider 路由) | ✅ | ✅ | 迁移 | +| 音色自动降级 | ✅ | ✅ | 迁移 | From 350177c7305fb2a372de2e092f50d6036ddd5f7f Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 10:30:16 +0800 Subject: [PATCH 02/55] =?UTF-8?q?feat(aibridge):=20=E9=98=B6=E6=AE=B50.1?= =?UTF-8?q?=20Cargo=20workspace=20=E9=AA=A8=E6=9E=B6?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - workspace 根 + 4 crate (aibridge-core/ffi/python/node) - aibridge-core 与 aibridge-ffi 编译通过 - rust-toolchain + .gitignore 更新 - 实现计划文档 --- .gitignore | 11 ++ Cargo.toml | 51 ++++++++ crates/aibridge-core/Cargo.toml | 22 ++++ crates/aibridge-core/src/lib.rs | 24 ++++ crates/aibridge-ffi/Cargo.toml | 24 ++++ crates/aibridge-ffi/src/lib.rs | 11 ++ crates/aibridge-node/Cargo.toml | 22 ++++ crates/aibridge-node/build.rs | 4 + crates/aibridge-node/src/lib.rs | 7 + crates/aibridge-python/Cargo.toml | 20 +++ crates/aibridge-python/src/lib.rs | 14 ++ ...2026-07-07-aibridge-implementation-plan.md | 122 ++++++++++++++++++ rust-toolchain.toml | 3 + 13 files changed, 335 insertions(+) create mode 100644 Cargo.toml create mode 100644 crates/aibridge-core/Cargo.toml create mode 100644 crates/aibridge-core/src/lib.rs create mode 100644 crates/aibridge-ffi/Cargo.toml create mode 100644 crates/aibridge-ffi/src/lib.rs create mode 100644 crates/aibridge-node/Cargo.toml create mode 100644 crates/aibridge-node/build.rs create mode 100644 crates/aibridge-node/src/lib.rs create mode 100644 crates/aibridge-python/Cargo.toml create mode 100644 crates/aibridge-python/src/lib.rs create mode 100644 docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md create mode 100644 rust-toolchain.toml diff --git a/.gitignore b/.gitignore index 9932a89..b5f618e 100644 --- a/.gitignore +++ b/.gitignore @@ -144,3 +144,14 @@ Thumbs.db # Project specific output/ *.egg-info/ + +# Rust(target/ 已在上方忽略) +# Cargo.lock 提交以保证可复现构建(workspace 含 cdylib 产物) + +# Node / napi-rs +node_modules/ +*.node + +# 原生库发布产物(target/ 外的) +*.dylib +*.dll diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..d91ee90 --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,51 @@ +[workspace] +resolver = "2" +members = [ + "crates/aibridge-core", + "crates/aibridge-ffi", + "crates/aibridge-python", + "crates/aibridge-node", +] + +[workspace.package] +version = "2.0.0-alpha.1" +edition = "2021" +rust-version = "1.75" +license = "MIT" +authors = ["WingkySky "] +repository = "https://github.com/WingkySky/aibridge" +description = "AIBridge - 多模态 AI 统一接口 SDK(一套 API 调用所有 AI 模型)" + +[workspace.dependencies] +# 内部 crate +aibridge-core = { path = "crates/aibridge-core" } +aibridge-ffi = { path = "crates/aibridge-ffi" } + +# 异步运行时 +tokio = { version = "1", features = ["full"] } +async-trait = "0.1" +async-stream = "0.3" +futures = "0.3" + +# HTTP 客户端 +reqwest = { version = "0.12", default-features = false, features = ["json", "stream", "http2", "rustls-tls"] } + +# 序列化 +serde = { version = "1", features = ["derive"] } +serde_json = "1" + +# 错误处理 +thiserror = "1" + +# 日志追踪 +tracing = "0.1" + +# 工具 +once_cell = "1" +bytes = "1" + +[profile.release] +opt-level = 3 +lto = true +codegen-units = 1 +strip = true diff --git a/crates/aibridge-core/Cargo.toml b/crates/aibridge-core/Cargo.toml new file mode 100644 index 0000000..91764c9 --- /dev/null +++ b/crates/aibridge-core/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "aibridge-core" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +authors.workspace = true +repository.workspace = true +description = "AIBridge 核心 - 多模态 AI 统一接口(Rust 核心,纯逻辑无 FFI)" + +[dependencies] +tokio.workspace = true +async-trait.workspace = true +async-stream.workspace = true +futures.workspace = true +reqwest.workspace = true +serde.workspace = true +serde_json.workspace = true +thiserror.workspace = true +tracing.workspace = true +once_cell.workspace = true +bytes.workspace = true diff --git a/crates/aibridge-core/src/lib.rs b/crates/aibridge-core/src/lib.rs new file mode 100644 index 0000000..38a44c3 --- /dev/null +++ b/crates/aibridge-core/src/lib.rs @@ -0,0 +1,24 @@ +//! AIBridge 核心 - 多模态 AI 统一接口 +//! +//! 一套 API 调用所有 AI 模型(chat / image / video / TTS / ASR / embed)。 +//! 本 crate 是纯 Rust 逻辑核心,不含 FFI 污染。 +//! Python/JS 通过 aibridge-python / aibridge-node 直连本 crate; +//! Go/JVM/.NET 通过 aibridge-ffi 的 C ABI 间接调用。 +//! +//! 对应 Python v1 (agn-sdk) 的 agn/ 目录,五层架构保持一致: +//! API 层(client) → 路由层(router) → 适配器层(adapter) → 核心层(http/retry/error/config) → 模型层(model) + +// 模块声明(阶段 0.2–0.4 逐步填充,暂注释) +// pub mod error; +// pub mod config; +// pub mod http; +// pub mod retry; +// pub mod util; +// pub mod model; +// pub mod adapter; +// pub mod adapters; +// pub mod client; +// pub mod router; + +/// crate 版本号 +pub const VERSION: &str = env!("CARGO_PKG_VERSION"); diff --git a/crates/aibridge-ffi/Cargo.toml b/crates/aibridge-ffi/Cargo.toml new file mode 100644 index 0000000..49d9364 --- /dev/null +++ b/crates/aibridge-ffi/Cargo.toml @@ -0,0 +1,24 @@ +[package] +name = "aibridge-ffi" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +authors.workspace = true +repository.workspace = true +description = "AIBridge C ABI - 供 Go/JVM/.NET 调用的 FFI 层" + +[lib] +# cdylib: 动态库(.so/.dylib/.dll);staticlib: 静态库;rlib: 内部测试用 +crate-type = ["cdylib", "staticlib", "rlib"] +name = "aibridge" + +[dependencies] +aibridge-core.workspace = true +tokio.workspace = true +once_cell.workspace = true +serde_json.workspace = true + +# cbindgen 头文件生成(阶段 0.5 启用) +# [build-dependencies] +# cbindgen = "0.27" diff --git a/crates/aibridge-ffi/src/lib.rs b/crates/aibridge-ffi/src/lib.rs new file mode 100644 index 0000000..e882359 --- /dev/null +++ b/crates/aibridge-ffi/src/lib.rs @@ -0,0 +1,11 @@ +//! AIBridge C ABI - FFI 层 +//! +//! 暴露 C ABI 供 Go/JVM/.NET 调用。 +//! Python/JS 通过 aibridge-python / aibridge-node 直连 aibridge-core,不走本层。 +//! +//! 设计要点(阶段 0.5 实现): +//! - 全局 tokio runtime(once_cell::Lazy),每个 FFI 调用 block_on +//! - 句柄式:aibridge_client_t / aibridge_stream_t(opaque) +//! - 复杂 struct 走 JSON 字符串边界,二进制走 aibridge_bytes_t +//! - 错误:aibridge_status_t 返回码 + aibridge_last_error() 线程局部槽 +//! - cbindgen 生成 include/aibridge.h diff --git a/crates/aibridge-node/Cargo.toml b/crates/aibridge-node/Cargo.toml new file mode 100644 index 0000000..5a73f2d --- /dev/null +++ b/crates/aibridge-node/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "aibridge-node" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +authors.workspace = true +repository.workspace = true +description = "AIBridge Node.js 绑定(napi-rs,直连 aibridge-core,原生 Promise)" + +[lib] +crate-type = ["cdylib"] + +[dependencies] +aibridge-core.workspace = true +napi = { version = "2", features = ["napi8", "tokio_rt"] } +napi-derive = "2" +tokio.workspace = true +serde_json.workspace = true + +[build-dependencies] +napi-build = "2" diff --git a/crates/aibridge-node/build.rs b/crates/aibridge-node/build.rs new file mode 100644 index 0000000..f596a7f --- /dev/null +++ b/crates/aibridge-node/build.rs @@ -0,0 +1,4 @@ +//! napi-rs 构建脚本 +fn main() { + napi_build::setup(); +} diff --git a/crates/aibridge-node/src/lib.rs b/crates/aibridge-node/src/lib.rs new file mode 100644 index 0000000..ddf4daa --- /dev/null +++ b/crates/aibridge-node/src/lib.rs @@ -0,0 +1,7 @@ +//! AIBridge Node.js 绑定(napi-rs) +//! +//! 直连 aibridge-core,原生 Promise / AsyncIterable 流式。 +//! 由 napi-rs 构建为 npm 包 `aibridge`。 +//! 阶段 0.6 填充 Client/chat/流式/错误映射。 + +// 阶段 0.6 填充:#[napi] 标注的 Client 及方法 diff --git a/crates/aibridge-python/Cargo.toml b/crates/aibridge-python/Cargo.toml new file mode 100644 index 0000000..d23cdc2 --- /dev/null +++ b/crates/aibridge-python/Cargo.toml @@ -0,0 +1,20 @@ +[package] +name = "aibridge-python" +version.workspace = true +edition.workspace = true +rust-version.workspace = true +license.workspace = true +authors.workspace = true +repository.workspace = true +description = "AIBridge Python 绑定(PyO3,直连 aibridge-core,原生 asyncio)" + +[lib] +crate-type = ["cdylib"] +name = "aibridge" + +[dependencies] +aibridge-core.workspace = true +# PyO3 0.22+ 原生支持 async fn,不再需要 pyo3-asyncio +pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310"] } +tokio.workspace = true +serde_json.workspace = true diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs new file mode 100644 index 0000000..c1f498c --- /dev/null +++ b/crates/aibridge-python/src/lib.rs @@ -0,0 +1,14 @@ +//! AIBridge Python 绑定(PyO3) +//! +//! 直连 aibridge-core,原生 asyncio 协程与 AsyncIterator 流式。 +//! 由 maturin 构建为 PyPI 包 `aibridge`。 +//! 阶段 0.6 填充 Client/chat/流式/错误映射。 + +use pyo3::prelude::*; + +/// Python 模块入口:import aibridge +#[pymodule] +fn aibridge(_py: Python, _m: &Bound) -> PyResult<()> { + // 阶段 0.6 填充:注册 Client / Router / 错误类 / 数据模型 + Ok(()) +} diff --git a/docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md b/docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md new file mode 100644 index 0000000..c17f5ee --- /dev/null +++ b/docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md @@ -0,0 +1,122 @@ +# AIBridge 实现计划 + +> 日期:2026-07-07 +> 依据设计:[docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md](../superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md) +> 实施方式:Claude 全权委托;阶段 0 串行,阶段 1+ 多 agent 并行 + +--- + +## 实施总策略 + +- **阶段 0(地基)**:Claude 串行搭建,保证命名/依赖/配置一致性(并行易产生不一致) +- **阶段 1(MVP 四 provider)**:多 agent 并行(每 provider 一个 agent),先由 Claude 搭好 OpenAI 兼容地基 +- **阶段 2(剩余适配器)**:多 agent 并行,按批次 2a/2b/2c +- **五语言绑定**:aibridge-core 与 aibridge-ffi 稳定后,多 agent 并行(每语言一个 agent) +- 每个任务带明确验收标准,作为 agent 完成判据与多 agent 编排的契约 + +--- + +## 阶段 0:地基(串行,2–3 周) + +### 0.1 Cargo workspace 骨架 +- workspace 根 `Cargo.toml` + 4 crate:aibridge-core / aibridge-ffi / aibridge-python / aibridge-node +- 统一依赖版本(workspace.dependencies) +- rust-toolchain.toml + 更新 .gitignore +- **验收**:`cargo build -p aibridge-core -p aibridge-ffi` 通过 + +### 0.2 aibridge-core 基础设施 +- `error.rs`:AibridgeError 枚举(thiserror),对齐 Python v1 错误分类 +- `config.rs`:ClientOptions + 环境变量加载(dotenv) +- `http.rs`:reqwest 封装(h2、超时、连接池) +- `retry.rs`:重试机制(指数退避,对应 tenacity) +- `util.rs` +- **验收**:核心层单测通过 + +### 0.3 数据模型层 +- `model/{chat,image,video,audio,common,options}.rs` +- serde struct + Builder derive +- **验收**:模型序列化/反序列化测试通过,与 Python fixture 对照 + +### 0.4 Adapter trait + Client + Router +- `adapter::{Adapter trait, Capabilities, create_adapter}` +- `client::Client`、`router::Router` +- **验收**:trait 可实现,Client/Router 编译通过 + +### 0.5 aibridge-ffi C ABI 骨架 +- 全局 tokio runtime(Lazy) +- 句柄管理(client/stream opaque) +- cbindgen 配置 + 生成 `aibridge.h` +- 基本函数:client_new/destroy/start、chat、chat_stream、stream_next/destroy、last_error、string/bytes_free +- **验收**:`cargo build -p aibridge-ffi` 产 cdylib,aibridge.h 生成 + +### 0.6 跨语言管线验证(openai chat stub) +- aibridge-python:PyO3 模块 + Client.chat async + 流式 AsyncIterator +- aibridge-node:napi 模块 + Client.chat async + 流式 AsyncIterable +- bindings/go:CGO 调 ffi + chat(goroutine+channel) +- bindings/jvm:JNA 调 ffi + chat(CompletableFuture) +- bindings/dotnet:P/Invoke 调 ffi + chat(Task) +- **验收**:五语言 hello world(含 async + 流式 + 错误)跑通 ← **最大技术风险点** + +--- + +## 阶段 1:MVP 四 provider(多 agent 并行,3–4 周) + +### 1.0 OpenAI 兼容地基(Claude 串行) +- `adapters/openai_compat.rs`:共享 HTTP 请求构造/响应解析/参数映射 +- `OPENAI_COMPATIBLE_MAPPING` 常量 +- **验收**:可被子适配器复用 + +### 1.1–1.4 四 provider 适配器(每 provider 一个 agent 并行) +| 任务 | provider | 协议 | 能力 | +|---|---|---|---| +| 1.1 | openai | 复用 compat | chat/image/embed/list_models | +| 1.2 | agnes | 复用 compat | chat/image/video/list_models | +| 1.3 | volcengine_cv | 独立 | image/video | +| 1.4 | gemini | 独立 | chat/image/embed/list_models | +- **验收**:每适配器单测通过(共享 mock fixture) + +### 1.5 五语言绑定完善(多 agent 并行,每语言一个) +- 四 provider 能力暴露到各语言 +- 跨语言一致性测试 +- **验收**:五语言 × 四 provider 全跑通 ← **用户另一项目可接入里程碑 M1** + +--- + +## 阶段 2:剩余适配器(多 agent 并行批次,4–6 周) + +### 2a OpenAI 兼容族(6 agent 并行) +azure, aggregation_platforms, additional_models, more_models, emerging_models, chinese(兼容部分) + +### 2b 独立协议(5 agent 并行) +anthropic, stability, runway, pika, kling + +### 2c 音频(edge-tts/elevenlabs/cartesia/deepgram/assemblyai,二进制载荷) + +--- + +## 阶段 3:发布收尾(2–3 周) + +- 3.1 CI 矩阵(平台 × 语言,交叉编译) +- 3.2 各语言包发布(PyPI/npm/Maven/NuGet/Go module) +- 3.3 Python v1→v2 迁移指南 +- 3.4 文档网站 +- 3.5 旧版 v1 归档 + 打 v2.0.0 tag + +--- + +## 多 agent 编排原则 + +- **适配器迁移**:每 provider 一个 agent,共享 fixture 与 OpenAiCompatAdapter 基础 +- **语言绑定**:每语言一个 agent,依赖 core/ffi 稳定后启动 +- **跨语言一致性测试**:单独 agent 汇总 +- **依赖顺序**:阶段 0 串行;阶段 1.0 地基串行;1.1–1.4 适配器并行;1.5 绑定并行 +- **每个 agent 任务契约**:设计文档引用 + 验收标准 + fixture 路径 + 命名规范 + +--- + +## 里程碑 + +- **M0**:阶段 0 完成,五语言 hello world 跑通(技术风险解除) +- **M1**:阶段 1 完成,四 provider 五语言可用(用户另一项目可接入)⭐ +- **M2**:阶段 2 完成,全量适配器迁移 +- **M3**:阶段 3 完成,v2.0.0 正式发布 diff --git a/rust-toolchain.toml b/rust-toolchain.toml new file mode 100644 index 0000000..73cb934 --- /dev/null +++ b/rust-toolchain.toml @@ -0,0 +1,3 @@ +[toolchain] +channel = "stable" +components = ["rustfmt", "clippy"] From 1b7414e11ada6cf09683e051f4cbaa264a15ff29 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 10:31:19 +0800 Subject: [PATCH 03/55] =?UTF-8?q?chore:=20=E6=8F=90=E4=BA=A4=20Cargo.lock?= =?UTF-8?q?=20=E4=BF=9D=E8=AF=81=E5=8F=AF=E5=A4=8D=E7=8E=B0=E6=9E=84?= =?UTF-8?q?=E5=BB=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 1782 ++++++++++++++++++++++++++++++++++++++++++++++++++++ 1 file changed, 1782 insertions(+) create mode 100644 Cargo.lock diff --git a/Cargo.lock b/Cargo.lock new file mode 100644 index 0000000..0e9b04d --- /dev/null +++ b/Cargo.lock @@ -0,0 +1,1782 @@ +# This file is automatically @generated by Cargo. +# It is not intended for manual editing. +version = 3 + +[[package]] +name = "aho-corasick" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ddd31a130427c27518df266943a5308ed92d4b226cc639f5a8f1002816174301" +dependencies = [ + "memchr", +] + +[[package]] +name = "aibridge-core" +version = "2.0.0-alpha.1" +dependencies = [ + "async-stream", + "async-trait", + "bytes", + "futures", + "once_cell", + "reqwest", + "serde", + "serde_json", + "thiserror 1.0.69", + "tokio", + "tracing", +] + +[[package]] +name = "aibridge-ffi" +version = "2.0.0-alpha.1" +dependencies = [ + "aibridge-core", + "once_cell", + "serde_json", + "tokio", +] + +[[package]] +name = "aibridge-node" +version = "2.0.0-alpha.1" +dependencies = [ + "aibridge-core", + "napi", + "napi-build", + "napi-derive", + "serde_json", + "tokio", +] + +[[package]] +name = "aibridge-python" +version = "2.0.0-alpha.1" +dependencies = [ + "aibridge-core", + "pyo3", + "serde_json", + "tokio", +] + +[[package]] +name = "async-stream" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b5a71a6f37880a80d1d7f19efd781e4b5de42c88f0722cc13bcb6cc2cfe8476" +dependencies = [ + "async-stream-impl", + "futures-core", + "pin-project-lite", +] + +[[package]] +name = "async-stream-impl" +version = "0.3.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7c24de15d275a1ecfd47a380fb4d5ec9bfe0933f309ed5e705b775596a3574d" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "async-trait" +version = "0.1.89" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9035ad2d096bed7955a320ee7e2230574d28fd3c3a0f186cbea1ff3c7eed5dbb" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "atomic-waker" +version = "1.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1505bd5d3d116872e7271a6d4e16d81d0c8570876c8de68093a09ac269d8aac0" + +[[package]] +name = "base64" +version = "0.22.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72b3254f16251a8381aa12e40e3c4d2f0199f8c6508fbecb9d91f575e0fbb8c6" + +[[package]] +name = "bitflags" +version = "2.13.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" + +[[package]] +name = "bumpalo" +version = "3.20.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" + +[[package]] +name = "bytes" +version = "1.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" + +[[package]] +name = "cc" +version = "1.2.66" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f5d6cac793997bd970000024b2934968efe83b382de4fdcf4fcb46b6ee4ad996" +dependencies = [ + "find-msvc-tools", + "shlex", +] + +[[package]] +name = "cfg-if" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9330f8b2ff13f34540b44e946ef35111825727b38d33286ef986142615121801" + +[[package]] +name = "cfg_aliases" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "613afe47fcd5fac7ccf1db93babcb082c5994d996f20b8b159f2ad1658eb5724" + +[[package]] +name = "chacha20" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" +dependencies = [ + "cfg-if", + "cpufeatures", + "rand_core", +] + +[[package]] +name = "convert_case" +version = "0.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec182b0ca2f35d8fc196cf3404988fd8b8c739a4d270ff118a398feb0cbec1ca" +dependencies = [ + "unicode-segmentation", +] + +[[package]] +name = "cpufeatures" +version = "0.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b2a41393f66f16b0823bb79094d54ac5fbd34ab292ddafb9a0456ac9f87d201" +dependencies = [ + "libc", +] + +[[package]] +name = "ctor" +version = "0.2.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a2785755761f3ddc1492979ce1e48d2c00d09311c39e4466429188f3dd6501" +dependencies = [ + "quote", + "syn", +] + +[[package]] +name = "displaydoc" +version = "0.2.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ac70aa55017e108007fbaf5aa0f54b021c98f92ff8af59d42eda9da96e3dd4f" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "equivalent" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "877a4ace8713b0bcf2a4e7eec82529c029f1d0619886d18145fea96c3ffe5c0f" + +[[package]] +name = "errno" +version = "0.3.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "39cab71617ae0d63f51a36d69f866391735b51691dbda63cf6f96d042b63efeb" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "find-msvc-tools" +version = "0.1.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5baebc0774151f905a1a2cc41989300b1e6fbb29aff0ceffa1064fdd3088d582" + +[[package]] +name = "fnv" +version = "1.0.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3f9eec918d3f24069decb9af1554cad7c880e2da24a9afd88aca000531ab82c1" + +[[package]] +name = "form_urlencoded" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb4cb245038516f5f85277875cdaa4f7d2c9a0fa0468de06ed190163b1581fcf" +dependencies = [ + "percent-encoding", +] + +[[package]] +name = "futures" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8b147ee9d1f6d097cef9ce628cd2ee62288d963e16fb287bd9286455b241382d" +dependencies = [ + "futures-channel", + "futures-core", + "futures-executor", + "futures-io", + "futures-sink", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-channel" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "07bbe89c50d7a535e539b8c17bc0b49bdb77747034daa8087407d655f3f7cc1d" +dependencies = [ + "futures-core", + "futures-sink", +] + +[[package]] +name = "futures-core" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7e3450815272ef58cec6d564423f6e755e25379b217b0bc688e295ba24df6b1d" + +[[package]] +name = "futures-executor" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "baf29c38818342a3b26b5b923639e7b1f4a61fc5e76102d4b1981c6dc7a7579d" +dependencies = [ + "futures-core", + "futures-task", + "futures-util", +] + +[[package]] +name = "futures-io" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cecba35d7ad927e23624b22ad55235f2239cfa44fd10428eecbeba6d6a717718" + +[[package]] +name = "futures-macro" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e835b70203e41293343137df5c0664546da5745f82ec9b84d40be8336958447b" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "futures-sink" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c39754e157331b013978ec91992bde1ac089843443c49cbc7f46150b0fad0893" + +[[package]] +name = "futures-task" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "037711b3d59c33004d3856fbdc83b99d4ff37a24768fa1be9ce3538a1cde4393" + +[[package]] +name = "futures-util" +version = "0.3.32" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "389ca41296e6190b48053de0321d02a77f32f8a5d2461dd38762c0593805c6d6" +dependencies = [ + "futures-channel", + "futures-core", + "futures-io", + "futures-macro", + "futures-sink", + "futures-task", + "memchr", + "pin-project-lite", + "slab", +] + +[[package]] +name = "getrandom" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff2abc00be7fca6ebc474524697ae276ad847ad0a6b3faa4bcb027e9a4614ad0" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "wasi", + "wasm-bindgen", +] + +[[package]] +name = "getrandom" +version = "0.4.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "300e883d756b2e4ec94e02791f39b04b522276138852cfc41d9fb7e904106099" +dependencies = [ + "cfg-if", + "js-sys", + "libc", + "r-efi", + "rand_core", + "wasm-bindgen", +] + +[[package]] +name = "h2" +version = "0.4.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6cb093c84e8bd9b188d4c4a8cb6579fc016968d14c99882163cd3ff402a4f155" +dependencies = [ + "atomic-waker", + "bytes", + "fnv", + "futures-core", + "futures-sink", + "http", + "indexmap", + "slab", + "tokio", + "tokio-util", + "tracing", +] + +[[package]] +name = "hashbrown" +version = "0.17.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed5909b6e89a2db4456e54cd5f673791d7eca6732202bbf2a9cc504fe2f9b84a" + +[[package]] +name = "heck" +version = "0.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" + +[[package]] +name = "http" +version = "1.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6970f50e31d6fc17d3fa27329444bfa74e196cf62e95052a3f6fee181dba6425" +dependencies = [ + "bytes", + "itoa", +] + +[[package]] +name = "http-body" +version = "1.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1efedce1fb8e6913f23e0c92de8e62cd5b772a67e7b3946df930a62566c93184" +dependencies = [ + "bytes", + "http", +] + +[[package]] +name = "http-body-util" +version = "0.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b021d93e26becf5dc7e1b75b1bed1fd93124b374ceb73f43d4d4eafec896a64a" +dependencies = [ + "bytes", + "futures-core", + "http", + "http-body", + "pin-project-lite", +] + +[[package]] +name = "httparse" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" + +[[package]] +name = "hyper" +version = "1.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "55281c53a1894c864990125767da440a4e630446785086f52523b20033b74498" +dependencies = [ + "atomic-waker", + "bytes", + "futures-channel", + "futures-core", + "h2", + "http", + "http-body", + "httparse", + "itoa", + "pin-project-lite", + "smallvec", + "tokio", + "want", +] + +[[package]] +name = "hyper-rustls" +version = "0.27.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "33ca68d021ef39cf6463ab54c1d0f5daf03377b70561305bb89a8f83aab66e0f" +dependencies = [ + "http", + "hyper", + "hyper-util", + "rustls", + "tokio", + "tokio-rustls", + "tower-service", + "webpki-roots", +] + +[[package]] +name = "hyper-util" +version = "0.1.20" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "96547c2556ec9d12fb1578c4eaf448b04993e7fb79cbaad930a656880a6bdfa0" +dependencies = [ + "base64", + "bytes", + "futures-channel", + "futures-util", + "http", + "http-body", + "hyper", + "ipnet", + "libc", + "percent-encoding", + "pin-project-lite", + "socket2", + "tokio", + "tower-service", + "tracing", +] + +[[package]] +name = "icu_collections" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2984d1cd16c883d7935b9e07e44071dca8d917fd52ecc02c04d5fa0b5a3f191c" +dependencies = [ + "displaydoc", + "potential_utf", + "utf8_iter", + "yoke", + "zerofrom", + "zerovec", +] + +[[package]] +name = "icu_locale_core" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92219b62b3e2b4d88ac5119f8904c10f8f61bf7e95b640d25ba3075e6cac2c29" +dependencies = [ + "displaydoc", + "litemap", + "tinystr", + "writeable", + "zerovec", +] + +[[package]] +name = "icu_normalizer" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c56e5ee99d6e3d33bd91c5d85458b6005a22140021cc324cea84dd0e72cff3b4" +dependencies = [ + "icu_collections", + "icu_normalizer_data", + "icu_properties", + "icu_provider", + "smallvec", + "zerovec", +] + +[[package]] +name = "icu_normalizer_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "da3be0ae77ea334f4da67c12f149704f19f81d1adf7c51cf482943e84a2bad38" + +[[package]] +name = "icu_properties" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bee3b67d0ea5c2cca5003417989af8996f8604e34fb9ddf96208a033901e70de" +dependencies = [ + "icu_collections", + "icu_locale_core", + "icu_properties_data", + "icu_provider", + "zerotrie", + "zerovec", +] + +[[package]] +name = "icu_properties_data" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e2bbb201e0c04f7b4b3e14382af113e17ba4f63e2c9d2ee626b720cbce54a14" + +[[package]] +name = "icu_provider" +version = "2.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "139c4cf31c8b5f33d7e199446eff9c1e02decfc2f0eec2c8d71f65befa45b421" +dependencies = [ + "displaydoc", + "icu_locale_core", + "writeable", + "yoke", + "zerofrom", + "zerotrie", + "zerovec", +] + +[[package]] +name = "idna" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3b0875f23caa03898994f6ddc501886a45c7d3d62d04d2d90788d47be1b1e4de" +dependencies = [ + "idna_adapter", + "smallvec", + "utf8_iter", +] + +[[package]] +name = "idna_adapter" +version = "1.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cb68373c0d6620ef8105e855e7745e18b0d00d3bdb07fb532e434244cdb9a714" +dependencies = [ + "icu_normalizer", + "icu_properties", +] + +[[package]] +name = "indexmap" +version = "2.14.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d466e9454f08e4a911e14806c24e16fba1b4c121d1ea474396f396069cf949d9" +dependencies = [ + "equivalent", + "hashbrown", +] + +[[package]] +name = "ipnet" +version = "2.12.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" + +[[package]] +name = "itoa" +version = "1.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8f42a60cbdf9a97f5d2305f08a87dc4e09308d1276d28c869c684d7777685682" + +[[package]] +name = "js-sys" +version = "0.3.103" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "53b44bfcdb3f8d5837a46dae1ca9660a837176eee74a28b229bc626816589102" +dependencies = [ + "cfg-if", + "futures-util", + "wasm-bindgen", +] + +[[package]] +name = "libc" +version = "0.2.186" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "68ab91017fe16c622486840e4c83c9a37afeff978bd239b5293d61ece587de66" + +[[package]] +name = "libloading" +version = "0.8.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d7c4b02199fee7c5d21a5ae7d8cfa79a6ef5bb2fc834d6e9058e89c825efdc55" +dependencies = [ + "cfg-if", + "windows-link", +] + +[[package]] +name = "litemap" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92daf443525c4cce67b150400bc2316076100ce0b3686209eb8cf3c31612e6f0" + +[[package]] +name = "lock_api" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "224399e74b87b5f3557511d98dff8b14089b3dadafcab6bb93eab67d3aace965" +dependencies = [ + "scopeguard", +] + +[[package]] +name = "log" +version = "0.4.33" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ceec5bc11778974d1bcb055b18002eba7f4b3518b6a0081b3af5f21666da9ad" + +[[package]] +name = "lru-slab" +version = "0.1.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "112b39cec0b298b6c1999fee3e31427f74f676e4cb9879ed1a121b43661a4154" + +[[package]] +name = "memchr" +version = "2.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "88904434abc2901f197fe8cc55f0445e7ded921dba5911dad2e2b39b48e663c4" + +[[package]] +name = "mio" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "02bd0af71c67b473010cbbc60715ee815645a4dc942899111f494b4b737d6fda" +dependencies = [ + "libc", + "wasi", + "windows-sys 0.61.2", +] + +[[package]] +name = "napi" +version = "2.16.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "55740c4ae1d8696773c78fdafd5d0e5fe9bc9f1b071c7ba493ba5c413a9184f3" +dependencies = [ + "bitflags", + "ctor", + "napi-derive", + "napi-sys", + "once_cell", + "tokio", +] + +[[package]] +name = "napi-build" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c9c366d2c8c60b86fa632df75f745509b52f9128f91a6bad4c796e44abb505e1" + +[[package]] +name = "napi-derive" +version = "2.16.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7cbe2585d8ac223f7d34f13701434b9d5f4eb9c332cccce8dee57ea18ab8ab0c" +dependencies = [ + "cfg-if", + "convert_case", + "napi-derive-backend", + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "napi-derive-backend" +version = "1.0.75" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1639aaa9eeb76e91c6ae66da8ce3e89e921cd3885e99ec85f4abacae72fc91bf" +dependencies = [ + "convert_case", + "once_cell", + "proc-macro2", + "quote", + "regex", + "semver", + "syn", +] + +[[package]] +name = "napi-sys" +version = "2.4.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "427802e8ec3a734331fec1035594a210ce1ff4dc5bc1950530920ab717964ea3" +dependencies = [ + "libloading", +] + +[[package]] +name = "once_cell" +version = "1.21.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" + +[[package]] +name = "parking_lot" +version = "0.12.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "93857453250e3077bd71ff98b6a65ea6621a19bb0f559a85248955ac12c45a1a" +dependencies = [ + "lock_api", + "parking_lot_core", +] + +[[package]] +name = "parking_lot_core" +version = "0.9.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2621685985a2ebf1c516881c026032ac7deafcda1a2c9b7850dc81e3dfcb64c1" +dependencies = [ + "cfg-if", + "libc", + "redox_syscall", + "smallvec", + "windows-link", +] + +[[package]] +name = "percent-encoding" +version = "2.3.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b4f627cb1b25917193a259e49bdad08f671f8d9708acfd5fe0a8c1455d87220" + +[[package]] +name = "pin-project-lite" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a89322df9ebe1c1578d689c92318e070967d1042b512afbe49518723f4e6d5cd" + +[[package]] +name = "portable-atomic" +version = "1.13.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c33a9471896f1c69cecef8d20cbe2f7accd12527ce60845ff44c153bb2a21b49" + +[[package]] +name = "potential_utf" +version = "0.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0103b1cef7ec0cf76490e969665504990193874ea05c85ff9bab8b911d0a0564" +dependencies = [ + "zerovec", +] + +[[package]] +name = "proc-macro2" +version = "1.0.106" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fd00f0bb2e90d81d1044c2b32617f68fcb9fa3bb7640c23e9c748e53fb30934" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "pyo3" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "91fd8e38a3b50ed1167fb981cd6fd60147e091784c427b8f7183a7ee32c31c12" +dependencies = [ + "libc", + "once_cell", + "portable-atomic", + "pyo3-build-config", + "pyo3-ffi", + "pyo3-macros", +] + +[[package]] +name = "pyo3-build-config" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e368e7ddfdeb98c9bca7f8383be1648fd84ab466bf2bc015e94008db6d35611e" +dependencies = [ + "target-lexicon", +] + +[[package]] +name = "pyo3-ffi" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7f29e10af80b1f7ccaf7f69eace800a03ecd13e883acfacc1e5d0988605f651e" +dependencies = [ + "libc", + "pyo3-build-config", +] + +[[package]] +name = "pyo3-macros" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df6e520eff47c45997d2fc7dd8214b25dd1310918bbb2642156ef66a67f29813" +dependencies = [ + "proc-macro2", + "pyo3-macros-backend", + "quote", + "syn", +] + +[[package]] +name = "pyo3-macros-backend" +version = "0.28.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4cdc218d835738f81c2338f822078af45b4afdf8b2e33cbb5916f108b813acb" +dependencies = [ + "heck", + "proc-macro2", + "pyo3-build-config", + "quote", + "syn", +] + +[[package]] +name = "quinn" +version = "0.11.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c1a41e437b6bbd489372cd4971de128e85c855f56c57f283d20ff016cf7c0a8" +dependencies = [ + "bytes", + "cfg_aliases", + "pin-project-lite", + "quinn-proto", + "quinn-udp", + "rustc-hash", + "rustls", + "socket2", + "thiserror 2.0.18", + "tokio", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-proto" +version = "0.11.16" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2f4bfc015262b9df63c8845072ce59068853ff5872180c2ce2f13038b970e560" +dependencies = [ + "bytes", + "getrandom 0.4.3", + "lru-slab", + "rand", + "rand_pcg", + "ring", + "rustc-hash", + "rustls", + "rustls-pki-types", + "slab", + "thiserror 2.0.18", + "tinyvec", + "tracing", + "web-time", +] + +[[package]] +name = "quinn-udp" +version = "0.5.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "35a133f956daabe89a61a685c2649f13d82d5aa4bd5d12d1277e1072a21c0694" +dependencies = [ + "cfg_aliases", + "libc", + "once_cell", + "socket2", + "tracing", + "windows-sys 0.61.2", +] + +[[package]] +name = "quote" +version = "1.0.46" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dfbc457d0c7a0759a614551b11a6409e5951f6c7537be1f1b7682b9ae9230368" +dependencies = [ + "proc-macro2", +] + +[[package]] +name = "r-efi" +version = "6.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" + +[[package]] +name = "rand" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" +dependencies = [ + "chacha20", + "getrandom 0.4.3", + "rand_core", +] + +[[package]] +name = "rand_core" +version = "0.10.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63b8176103e19a2643978565ca18b50549f6101881c443590420e4dc998a3c69" + +[[package]] +name = "rand_pcg" +version = "0.10.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" +dependencies = [ + "rand_core", +] + +[[package]] +name = "redox_syscall" +version = "0.5.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ed2bf2547551a7053d6fdfafda3f938979645c44812fbfcda098faae3f1a362d" +dependencies = [ + "bitflags", +] + +[[package]] +name = "regex" +version = "1.12.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f1292b7759ae1cb9ec195452d1390a074f0cd8541ab7a5a8c31cd6db45d4a6ba" +dependencies = [ + "aho-corasick", + "memchr", + "regex-automata", + "regex-syntax", +] + +[[package]] +name = "regex-automata" +version = "0.4.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6e1dd4122fc1595e8162618945476892eefca7b88c52820e74af6262213cae8f" +dependencies = [ + "aho-corasick", + "memchr", + "regex-syntax", +] + +[[package]] +name = "regex-syntax" +version = "0.8.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d6f6ff9a378485b298a5286656da665ba74413d36db0979633275d2e708145d4" + +[[package]] +name = "reqwest" +version = "0.12.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "eddd3ca559203180a307f12d114c268abf583f59b03cb906fd0b3ff8646c1147" +dependencies = [ + "base64", + "bytes", + "futures-core", + "futures-util", + "h2", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-rustls", + "hyper-util", + "js-sys", + "log", + "percent-encoding", + "pin-project-lite", + "quinn", + "rustls", + "rustls-pki-types", + "serde", + "serde_json", + "serde_urlencoded", + "sync_wrapper", + "tokio", + "tokio-rustls", + "tokio-util", + "tower", + "tower-http", + "tower-service", + "url", + "wasm-bindgen", + "wasm-bindgen-futures", + "wasm-streams", + "web-sys", + "webpki-roots", +] + +[[package]] +name = "ring" +version = "0.17.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4689e6c2294d81e88dc6261c768b63bc4fcdb852be6d1352498b114f61383b7" +dependencies = [ + "cc", + "cfg-if", + "getrandom 0.2.17", + "libc", + "untrusted", + "windows-sys 0.52.0", +] + +[[package]] +name = "rustc-hash" +version = "2.1.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" + +[[package]] +name = "rustls" +version = "0.23.41" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6b92b125634d9b795e7beca796cc790df15a7fb38323bf3196fda83292d06b1f" +dependencies = [ + "once_cell", + "ring", + "rustls-pki-types", + "rustls-webpki", + "subtle", + "zeroize", +] + +[[package]] +name = "rustls-pki-types" +version = "1.15.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "764899a24af3980067ee14bc143654f297b22eaebfe3c7b6b211920a5a59b046" +dependencies = [ + "web-time", + "zeroize", +] + +[[package]] +name = "rustls-webpki" +version = "0.103.13" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "61c429a8649f110dddef65e2a5ad240f747e85f7758a6bccc7e5777bd33f756e" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + +[[package]] +name = "rustversion" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf54715a573b99ac80df0bc206da022bcd442c974952c7b9720069370852e21f" + +[[package]] +name = "ryu" +version = "1.0.23" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9774ba4a74de5f7b1c1451ed6cd5285a32eddb5cccb8cc655a4e50009e06477f" + +[[package]] +name = "scopeguard" +version = "1.2.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "94143f37725109f92c262ed2cf5e59bce7498c01bcc1502d7b9afe439a4e9f49" + +[[package]] +name = "semver" +version = "1.0.28" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8a7852d02fc848982e0c167ef163aaff9cd91dc640ba85e263cb1ce46fae51cd" + +[[package]] +name = "serde" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9a8e94ea7f378bd32cbbd37198a4a91436180c5bb472411e48b5ec2e2124ae9e" +dependencies = [ + "serde_core", + "serde_derive", +] + +[[package]] +name = "serde_core" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "41d385c7d4ca58e59fc732af25c3983b67ac852c1a25000afe1175de458b67ad" +dependencies = [ + "serde_derive", +] + +[[package]] +name = "serde_derive" +version = "1.0.228" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d540f220d3187173da220f885ab66608367b6574e925011a9353e4badda91d79" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "serde_json" +version = "1.0.150" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e8014e44b4736ed0538adeecded0fce2a272f22dc9578a7eb6b2d9993c74cfb9" +dependencies = [ + "itoa", + "memchr", + "serde", + "serde_core", + "zmij", +] + +[[package]] +name = "serde_urlencoded" +version = "0.7.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3491c14715ca2294c4d6a88f15e84739788c1d030eed8c110436aafdaa2f3fd" +dependencies = [ + "form_urlencoded", + "itoa", + "ryu", + "serde", +] + +[[package]] +name = "shlex" +version = "2.0.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f8fadd59c855ef2080decdef8ff161eb6661b86933c9d82e5ba29dc602a55aba" + +[[package]] +name = "signal-hook-registry" +version = "1.4.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c4db69cba1110affc0e9f7bcd48bbf87b3f4fc7c61fc9155afd4c469eb3d6c1b" +dependencies = [ + "errno", + "libc", +] + +[[package]] +name = "slab" +version = "0.4.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0c790de23124f9ab44544d7ac05d60440adc586479ce501c1d6d7da3cd8c9cf5" + +[[package]] +name = "smallvec" +version = "1.15.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ed6a63f02c8539c91a8685a86f4099661ba3da017932f6ebbea6de3f0fa7c90" + +[[package]] +name = "socket2" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52d1cfed4120b4d927bf7c0f86d2087a4a7d6027c906d9f9d525a80573b9be51" +dependencies = [ + "libc", + "windows-sys 0.61.2", +] + +[[package]] +name = "stable_deref_trait" +version = "1.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" + +[[package]] +name = "subtle" +version = "2.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "13c2bddecc57b384dee18652358fb23172facb8a2c51ccc10d74c157bdea3292" + +[[package]] +name = "syn" +version = "2.0.118" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1b9ae57f904213ebb649ce6895b8a66c66f0203b9319718f69a5612a065b1422" +dependencies = [ + "proc-macro2", + "quote", + "unicode-ident", +] + +[[package]] +name = "sync_wrapper" +version = "1.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0bf256ce5efdfa370213c1dabab5935a12e49f2c58d15e9eac2870d3b4f27263" +dependencies = [ + "futures-core", +] + +[[package]] +name = "synstructure" +version = "0.13.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "728a70f3dbaf5bab7f0c4b1ac8d7ae5ea60a4b5549c8a5914361c99147a709d2" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "target-lexicon" +version = "0.13.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca" + +[[package]] +name = "thiserror" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6aaf5339b578ea85b50e080feb250a3e8ae8cfcdff9a461c9ec2904bc923f52" +dependencies = [ + "thiserror-impl 1.0.69", +] + +[[package]] +name = "thiserror" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4288b5bcbc7920c07a1149a35cf9590a2aa808e0bc1eafaade0b80947865fbc4" +dependencies = [ + "thiserror-impl 2.0.18", +] + +[[package]] +name = "thiserror-impl" +version = "1.0.69" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4fee6c4efc90059e10f81e6d42c60a18f76588c3d74cb83a0b242a2b6c7504c1" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "thiserror-impl" +version = "2.0.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebc4ee7f67670e9b64d05fa4253e753e016c6c95ff35b89b7941d6b856dec1d5" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tinystr" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8323304221c2a851516f22236c5722a72eaa19749016521d6dff0824447d96d" +dependencies = [ + "displaydoc", + "zerovec", +] + +[[package]] +name = "tinyvec" +version = "1.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3e61e67053d25a4e82c844e8424039d9745781b3fc4f32b8d55ed50f5f667ef3" +dependencies = [ + "tinyvec_macros", +] + +[[package]] +name = "tinyvec_macros" +version = "0.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1f3ccbac311fea05f86f61904b462b55fb3df8837a366dfc601a0161d0532f20" + +[[package]] +name = "tokio" +version = "1.52.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8fc7f01b389ac15039e4dc9531aa973a135d7a4135281b12d7c1bc79fd57fffe" +dependencies = [ + "bytes", + "libc", + "mio", + "parking_lot", + "pin-project-lite", + "signal-hook-registry", + "socket2", + "tokio-macros", + "windows-sys 0.61.2", +] + +[[package]] +name = "tokio-macros" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "385a6cb71ab9ab790c5fe8d67f1645e6c450a7ce006a33de03daa956cf70a496" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tokio-rustls" +version = "0.26.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" +dependencies = [ + "rustls", + "tokio", +] + +[[package]] +name = "tokio-util" +version = "0.7.18" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ae9cec805b01e8fc3fd2fe289f89149a9b66dd16786abd8b19cfa7b48cb0098" +dependencies = [ + "bytes", + "futures-core", + "futures-sink", + "pin-project-lite", + "tokio", +] + +[[package]] +name = "tower" +version = "0.5.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ebe5ef63511595f1344e2d5cfa636d973292adc0eec1f0ad45fae9f0851ab1d4" +dependencies = [ + "futures-core", + "futures-util", + "pin-project-lite", + "sync_wrapper", + "tokio", + "tower-layer", + "tower-service", +] + +[[package]] +name = "tower-http" +version = "0.6.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4cfcf7e2740e6fc6d4d688b4ef00650406bb94adf4731e43c096c3a19fe40840" +dependencies = [ + "bitflags", + "bytes", + "futures-util", + "http", + "http-body", + "pin-project-lite", + "tower", + "tower-layer", + "tower-service", + "url", +] + +[[package]] +name = "tower-layer" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "121c2a6cda46980bb0fcd1647ffaf6cd3fc79a013de288782836f6df9c48780e" + +[[package]] +name = "tower-service" +version = "0.3.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8df9b6e13f2d32c91b9bd719c00d1958837bc7dec474d94952798cc8e69eeec3" + +[[package]] +name = "tracing" +version = "0.1.44" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "63e71662fa4b2a2c3a26f570f037eb95bb1f85397f3cd8076caed2f026a6d100" +dependencies = [ + "pin-project-lite", + "tracing-attributes", + "tracing-core", +] + +[[package]] +name = "tracing-attributes" +version = "0.1.31" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7490cfa5ec963746568740651ac6781f701c9c5ea257c58e057f3ba8cf69e8da" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "tracing-core" +version = "0.1.36" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "db97caf9d906fbde555dd62fa95ddba9eecfd14cb388e4f491a66d74cd5fb79a" +dependencies = [ + "once_cell", +] + +[[package]] +name = "try-lock" +version = "0.2.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" + +[[package]] +name = "unicode-ident" +version = "1.0.24" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6e4313cd5fcd3dad5cafa179702e2b244f760991f45397d14d4ebf38247da75" + +[[package]] +name = "unicode-segmentation" +version = "1.13.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c6f5d3c3b1bf09027a88a6bc961fc00497d651009560b5463668dc81b0fa87a8" + +[[package]] +name = "untrusted" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8ecb6da28b8a351d773b68d5825ac39017e680750f980f3a1a85cd8dd28a47c1" + +[[package]] +name = "url" +version = "2.5.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ff67a8a4397373c3ef660812acab3268222035010ab8680ec4215f38ba3d0eed" +dependencies = [ + "form_urlencoded", + "idna", + "percent-encoding", + "serde", +] + +[[package]] +name = "utf8_iter" +version = "1.0.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" + +[[package]] +name = "want" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bfa7760aed19e106de2c7c0b581b509f2f25d3dacaf737cb82ac61bc6d760b0e" +dependencies = [ + "try-lock", +] + +[[package]] +name = "wasi" +version = "0.11.1+wasi-snapshot-preview1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" + +[[package]] +name = "wasm-bindgen" +version = "0.2.126" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4b067c0c11094aef6b7a801c1e34a26affafdf3d051dba08456b868789aaf9a4" +dependencies = [ + "cfg-if", + "once_cell", + "rustversion", + "wasm-bindgen-macro", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-futures" +version = "0.4.76" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c62df1340f32221cb9c54d6a27b030e3dba64361d4a95bed55f9aacb44da291d" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "wasm-bindgen-macro" +version = "0.2.126" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "167ce5e579f6bcf889c4f7175a8a5a585de84e8ff93976ce393efa5f2837aab1" +dependencies = [ + "quote", + "wasm-bindgen-macro-support", +] + +[[package]] +name = "wasm-bindgen-macro-support" +version = "0.2.126" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f3997c7839262f4ef12cf90b818d6340c18e80f263f1a94bf157d0ec4420380e" +dependencies = [ + "bumpalo", + "proc-macro2", + "quote", + "syn", + "wasm-bindgen-shared", +] + +[[package]] +name = "wasm-bindgen-shared" +version = "0.2.126" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "dc1b4cb0cc549fcf58d7dfc081778139b3d283a081644e833e84682ad71cea24" +dependencies = [ + "unicode-ident", +] + +[[package]] +name = "wasm-streams" +version = "0.4.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "15053d8d85c7eccdbefef60f06769760a563c7f0a9d6902a13d35c7800b0ad65" +dependencies = [ + "futures-util", + "js-sys", + "wasm-bindgen", + "wasm-bindgen-futures", + "web-sys", +] + +[[package]] +name = "web-sys" +version = "0.3.103" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8622dcb61c0bcc9fffa6938bed81210af2da9a7e4a1a834b2e37a59b6dfb6141" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "web-time" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5a6580f308b1fad9207618087a65c04e7a10bc77e02c8e84e9b00dd4b12fa0bb" +dependencies = [ + "js-sys", + "wasm-bindgen", +] + +[[package]] +name = "webpki-roots" +version = "1.0.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf85cb06032201fa7c6f829d7db5a7e5aa45bcc0655327713065f6f0576731bf" +dependencies = [ + "rustls-pki-types", +] + +[[package]] +name = "windows-link" +version = "0.2.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "f0805222e57f7521d6a62e36fa9163bc891acd422f971defe97d64e70d0a4fe5" + +[[package]] +name = "windows-sys" +version = "0.52.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "282be5f36a8ce781fad8c8ae18fa3f9beff57ec1b52cb3de0789201425d9a33d" +dependencies = [ + "windows-targets", +] + +[[package]] +name = "windows-sys" +version = "0.61.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ae137229bcbd6cdf0f7b80a31df61766145077ddf49416a728b02cb3921ff3fc" +dependencies = [ + "windows-link", +] + +[[package]] +name = "windows-targets" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9b724f72796e036ab90c1021d4780d4d3d648aca59e491e6b98e725b84e99973" +dependencies = [ + "windows_aarch64_gnullvm", + "windows_aarch64_msvc", + "windows_i686_gnu", + "windows_i686_gnullvm", + "windows_i686_msvc", + "windows_x86_64_gnu", + "windows_x86_64_gnullvm", + "windows_x86_64_msvc", +] + +[[package]] +name = "windows_aarch64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a4622180e7a0ec044bb555404c800bc9fd9ec262ec147edd5989ccd0c02cd3" + +[[package]] +name = "windows_aarch64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09ec2a7bb152e2252b53fa7803150007879548bc709c039df7627cabbd05d469" + +[[package]] +name = "windows_i686_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "8e9b5ad5ab802e97eb8e295ac6720e509ee4c243f69d781394014ebfe8bbfa0b" + +[[package]] +name = "windows_i686_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0eee52d38c090b3caa76c563b86c3a4bd71ef1a819287c19d586d7334ae8ed66" + +[[package]] +name = "windows_i686_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "240948bc05c5e7c6dabba28bf89d89ffce3e303022809e73deaefe4f6ec56c66" + +[[package]] +name = "windows_x86_64_gnu" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "147a5c80aabfbf0c7d901cb5895d1de30ef2907eb21fbbab29ca94c5b08b1a78" + +[[package]] +name = "windows_x86_64_gnullvm" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "24d5b23dc417412679681396f2b49f3de8c1473deb516bd34410872eff51ed0d" + +[[package]] +name = "windows_x86_64_msvc" +version = "0.52.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" + +[[package]] +name = "writeable" +version = "0.6.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ffae5123b2d3fc086436f8834ae3ab053a283cfac8fe0a0b8eaae044768a4c4" + +[[package]] +name = "yoke" +version = "0.8.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "709fe23a0424b6a435d82152b1bd3fdfb0833487d5fa90d05d42762a9891fef5" +dependencies = [ + "stable_deref_trait", + "yoke-derive", + "zerofrom", +] + +[[package]] +name = "yoke-derive" +version = "0.8.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "de844c262c8848816172cef550288e7dc6c7b7814b4ee56b3e1553f275f1858e" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zerofrom" +version = "0.1.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0ec05a11813ea801ff6d75110ad09cd0824ddba17dfe17128ea0d5f68e6c5272" +dependencies = [ + "zerofrom-derive", +] + +[[package]] +name = "zerofrom-derive" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "11532158c46691caf0f2593ea8358fed6bbf68a0315e80aae9bd41fbade684a1" +dependencies = [ + "proc-macro2", + "quote", + "syn", + "synstructure", +] + +[[package]] +name = "zeroize" +version = "1.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e13c156562582aa81c60cb29407084cdb54c4164760106ab78e6c5b0858cf64e" + +[[package]] +name = "zerotrie" +version = "0.2.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0f9152d31db0792fa83f70fb2f83148effb5c1f5b8c7686c3459e361d9bc20bf" +dependencies = [ + "displaydoc", + "yoke", + "zerofrom", +] + +[[package]] +name = "zerovec" +version = "0.11.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90f911cbc359ab6af17377d242225f4d75119aec87ea711a880987b18cd7b239" +dependencies = [ + "yoke", + "zerofrom", + "zerovec-derive", +] + +[[package]] +name = "zerovec-derive" +version = "0.11.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "625dc425cab0dca6dc3c3319506e6593dcb08a9f387ea3b284dbd52a92c40555" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + +[[package]] +name = "zmij" +version = "1.0.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b8848ee67ecc8aedbaf3e4122217aff892639231befc6a1b58d29fff4c2cabaa" From 414678d66da2e1871c7ff956881718a86fb5d0f4 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 11:01:15 +0800 Subject: [PATCH 04/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?0.2-0.4=20=E5=9F=BA=E7=A1=80=E8=AE=BE=E6=96=BD/=E6=95=B0?= =?UTF-8?q?=E6=8D=AE=E6=A8=A1=E5=9E=8B/Adapter/Client/Router?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 71 +- crates/aibridge-core/Cargo.toml | 4 + crates/aibridge-core/src/adapter/base.rs | 465 ++++++++++ crates/aibridge-core/src/adapter/factory.rs | 170 ++++ crates/aibridge-core/src/adapter/mod.rs | 15 + crates/aibridge-core/src/adapters/mod.rs | 11 + crates/aibridge-core/src/client.rs | 194 +++++ crates/aibridge-core/src/config.rs | 489 +++++++++++ crates/aibridge-core/src/error.rs | 409 +++++++++ crates/aibridge-core/src/http.rs | 284 +++++++ crates/aibridge-core/src/lib.rs | 22 +- crates/aibridge-core/src/model/audio.rs | 514 +++++++++++ crates/aibridge-core/src/model/chat.rs | 660 ++++++++++++++ crates/aibridge-core/src/model/common.rs | 385 +++++++++ crates/aibridge-core/src/model/image.rs | 382 +++++++++ crates/aibridge-core/src/model/mod.rs | 40 + crates/aibridge-core/src/model/options.rs | 416 +++++++++ crates/aibridge-core/src/model/video.rs | 385 +++++++++ crates/aibridge-core/src/retry.rs | 306 +++++++ crates/aibridge-core/src/router.rs | 899 ++++++++++++++++++++ crates/aibridge-core/src/util.rs | 316 +++++++ 21 files changed, 6421 insertions(+), 16 deletions(-) create mode 100644 crates/aibridge-core/src/adapter/base.rs create mode 100644 crates/aibridge-core/src/adapter/factory.rs create mode 100644 crates/aibridge-core/src/adapter/mod.rs create mode 100644 crates/aibridge-core/src/adapters/mod.rs create mode 100644 crates/aibridge-core/src/client.rs create mode 100644 crates/aibridge-core/src/config.rs create mode 100644 crates/aibridge-core/src/error.rs create mode 100644 crates/aibridge-core/src/http.rs create mode 100644 crates/aibridge-core/src/model/audio.rs create mode 100644 crates/aibridge-core/src/model/chat.rs create mode 100644 crates/aibridge-core/src/model/common.rs create mode 100644 crates/aibridge-core/src/model/image.rs create mode 100644 crates/aibridge-core/src/model/mod.rs create mode 100644 crates/aibridge-core/src/model/options.rs create mode 100644 crates/aibridge-core/src/model/video.rs create mode 100644 crates/aibridge-core/src/retry.rs create mode 100644 crates/aibridge-core/src/router.rs create mode 100644 crates/aibridge-core/src/util.rs diff --git a/Cargo.lock b/Cargo.lock index 0e9b04d..3f77fd2 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -17,9 +17,11 @@ version = "2.0.0-alpha.1" dependencies = [ "async-stream", "async-trait", + "base64", "bytes", "futures", "once_cell", + "rand 0.8.6", "reqwest", "serde", "serde_json", @@ -153,7 +155,7 @@ checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" dependencies = [ "cfg-if", "cpufeatures", - "rand_core", + "rand_core 0.10.1", ] [[package]] @@ -343,7 +345,7 @@ dependencies = [ "js-sys", "libc", "r-efi", - "rand_core", + "rand_core 0.10.1", "wasm-bindgen", ] @@ -787,6 +789,15 @@ dependencies = [ "zerovec", ] +[[package]] +name = "ppv-lite86" +version = "0.2.21" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85eae3c4ed2f50dcfe72643da4befc30deadb458a9b590d720cde2f2b1e97da9" +dependencies = [ + "zerocopy", +] + [[package]] name = "proc-macro2" version = "1.0.106" @@ -883,7 +894,7 @@ dependencies = [ "bytes", "getrandom 0.4.3", "lru-slab", - "rand", + "rand 0.10.2", "rand_pcg", "ring", "rustc-hash", @@ -925,6 +936,17 @@ version = "6.0.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "f8dcc9c7d52a811697d2151c701e0d08956f92b0e24136cf4cf27b57a6a0d9bf" +[[package]] +name = "rand" +version = "0.8.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" +dependencies = [ + "libc", + "rand_chacha", + "rand_core 0.6.4", +] + [[package]] name = "rand" version = "0.10.2" @@ -933,7 +955,26 @@ checksum = "c7f5fa3a058cd35567ef9bfa5e75732bee0f9e4c55fa90477bef2dfcdbc4be80" dependencies = [ "chacha20", "getrandom 0.4.3", - "rand_core", + "rand_core 0.10.1", +] + +[[package]] +name = "rand_chacha" +version = "0.3.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e6c10a63a0fa32252be49d21e7709d4d4baf8d231c2dbce1eaa8141b9b127d88" +dependencies = [ + "ppv-lite86", + "rand_core 0.6.4", +] + +[[package]] +name = "rand_core" +version = "0.6.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "ec0be4795e2f6a28069bec0b5ff3e2ac9bafc99e6a9a7dc3547996c5c816922c" +dependencies = [ + "getrandom 0.2.17", ] [[package]] @@ -948,7 +989,7 @@ version = "0.10.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "caa0f4137e1c0a72f4c651489402276c8e8e1cf081f3b0ba156d2cbeef09e86a" dependencies = [ - "rand_core", + "rand_core 0.10.1", ] [[package]] @@ -1715,6 +1756,26 @@ dependencies = [ "synstructure", ] +[[package]] +name = "zerocopy" +version = "0.8.53" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "75726053136156d419e285b9b7eddaaea9e3fea6ce32eed44a89901f0bd98de1" +dependencies = [ + "zerocopy-derive", +] + +[[package]] +name = "zerocopy-derive" +version = "0.8.53" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "4714fd92cf900833d49538023a9b3915155210801d1c1169eba513b2addefd71" +dependencies = [ + "proc-macro2", + "quote", + "syn", +] + [[package]] name = "zerofrom" version = "0.1.8" diff --git a/crates/aibridge-core/Cargo.toml b/crates/aibridge-core/Cargo.toml index 91764c9..896a933 100644 --- a/crates/aibridge-core/Cargo.toml +++ b/crates/aibridge-core/Cargo.toml @@ -20,3 +20,7 @@ thiserror.workspace = true tracing.workspace = true once_cell.workspace = true bytes.workspace = true +# base64 编解码(util.rs 用,处理 data URI / 音频 base64) +base64 = "0.22" +# 随机数(router.rs 的 round_robin/random/weighted 策略用) +rand = "0.8" diff --git a/crates/aibridge-core/src/adapter/base.rs b/crates/aibridge-core/src/adapter/base.rs new file mode 100644 index 0000000..90dc1db --- /dev/null +++ b/crates/aibridge-core/src/adapter/base.rs @@ -0,0 +1,465 @@ +//! Adapter trait 与能力定义 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/base.py`。 +//! +//! 设计要点(与设计文档 5.2 节一致): +//! - `Adapter` trait 用 `#[async_trait]` +//! - 不支持的方法默认返 `UnsupportedCapability`,trait 提供默认实现 +//! - `Capabilities` 为枚举(替代 Python 的字符串常量),更类型安全 + +use async_trait::async_trait; +use futures::stream::BoxStream; +use std::collections::HashSet; + +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; +use crate::model::{ + ChatCompletion, ChatCompletionChunk, ChatRequest, EmbedRequest, EmbeddingResult, ImageRequest, + ImageResult, SpeechRequest, SpeechResult, TranscribeRequest, TranscriptionResult, VideoRequest, + VideoStatus, VideoTask, +}; + +/// 流式对话的类型别名 +/// +/// 用 `BoxStream` 封装 `impl Stream>`, +/// 便于跨 `dyn Adapter` 使用(无法直接返回 `impl Stream`)。 +pub type ChatStream = BoxStream<'static, Result>; + +/// 能力枚举 +/// +/// 对应 Python v1 `Capabilities` 字符串常量类。 +/// 用枚举替代字符串,编译期保证能力名拼写正确。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash)] +pub enum Capabilities { + // 对话能力 + /// 文本对话 + Chat, + /// 流式文本对话 + ChatStream, + /// 视觉理解 + Vision, + /// 工具调用 + ToolCall, + /// 推理/思考模式 + Reasoning, + /// JSON 模式 + JsonMode, + /// 联网搜索 + WebSearch, + + // 图像能力 + /// 图像生成 + ImageGenerate, + /// 图像编辑(图生图) + ImageEdit, + + // 视频能力 + /// 视频生成 + VideoGenerate, + /// 文生视频 + VideoText2Video, + /// 图生视频 + VideoImage2Video, + + // 嵌入能力 + /// 文本嵌入 + Embedding, + + // 音频能力 + /// 语音转文字 + AudioTranscribe, + /// 文字转语音 + AudioSpeech, + /// 音色查询 + ListVoices, +} + +impl Capabilities { + /// 转为字符串标识(用于序列化/对照 Python v1) + pub fn as_str(&self) -> &'static str { + match self { + Self::Chat => "chat", + Self::ChatStream => "chat_stream", + Self::Vision => "vision", + Self::ToolCall => "tool_call", + Self::Reasoning => "reasoning", + Self::JsonMode => "json_mode", + Self::WebSearch => "web_search", + Self::ImageGenerate => "image", + Self::ImageEdit => "image_edit", + Self::VideoGenerate => "video", + Self::VideoText2Video => "text2video", + Self::VideoImage2Video => "image2video", + Self::Embedding => "embedding", + Self::AudioTranscribe => "audio_transcribe", + Self::AudioSpeech => "audio_speech", + Self::ListVoices => "list_voices", + } + } +} + +/// 能力集合(用 HashSet 便于快速查找) +pub type CapabilitySet = HashSet; + +/// Adapter trait +/// +/// 所有 Provider 适配器必须实现此 trait。 +/// 适配器负责将统一接口转换为各 Provider 的特定 API 调用格式, +/// 并将响应归一化为统一的数据结构。 +/// +/// 不支持的方法默认返 `UnsupportedCapability`,子类按需 override。 +#[async_trait] +pub trait Adapter: Send + Sync { + /// Provider 类型标识(如 "agnes"、"openai") + fn provider_type(&self) -> &str; + + /// Provider 显示名称 + fn provider_name(&self) -> &str; + + /// 支持的能力集合 + fn capabilities(&self) -> CapabilitySet; + + /// 是否需要 API Key(免费 Provider 如 Edge TTS 返回 false) + fn requires_api_key(&self) -> bool { + true + } + + /// 启动适配器(初始化 HTTP 客户端、连接池等资源) + async fn start(&mut self) -> Result<()>; + + /// 关闭适配器(释放所有资源) + async fn close(&mut self) -> Result<()>; + + /// 文本对话 + async fn chat(&self, req: ChatRequest) -> Result { + let _ = req; + Err(unsupported(self.provider_type(), Capabilities::Chat)) + } + + /// 流式文本对话 + /// + /// 返回 `BoxStream>`。 + /// 默认实现返 `UnsupportedCapability`。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + let _ = req; + Err(unsupported(self.provider_type(), Capabilities::ChatStream)) + } + + /// 图像生成 + async fn image_generate(&self, req: ImageRequest) -> Result { + let _ = req; + Err(unsupported( + self.provider_type(), + Capabilities::ImageGenerate, + )) + } + + /// 创建视频生成任务 + async fn video_create(&self, req: VideoRequest) -> Result { + let _ = req; + Err(unsupported( + self.provider_type(), + Capabilities::VideoGenerate, + )) + } + + /// 查询视频任务状态 + async fn video_poll(&self, task_id: &str, model: &str) -> Result { + let _ = task_id; + let _ = model; + Err(unsupported( + self.provider_type(), + Capabilities::VideoGenerate, + )) + } + + /// 文本嵌入 + async fn embed(&self, req: EmbedRequest) -> Result { + let _ = req; + Err(unsupported(self.provider_type(), Capabilities::Embedding)) + } + + /// 语音转文字 + async fn transcribe(&self, req: TranscribeRequest) -> Result { + let _ = req; + Err(unsupported( + self.provider_type(), + Capabilities::AudioTranscribe, + )) + } + + /// 文字转语音 + async fn speech(&self, req: SpeechRequest) -> Result { + let _ = req; + Err(unsupported(self.provider_type(), Capabilities::AudioSpeech)) + } + + /// 获取可用模型列表(实时拉取) + async fn list_models(&self, filter: Option) -> Result>; + + /// 列出可用音色 + async fn list_voices(&self, language: Option<&str>) -> Result> { + let _ = language; + Err(unsupported(self.provider_type(), Capabilities::ListVoices)) + } + + /// 推荐可用音色(按语言/性别过滤) + async fn recommend_voices( + &self, + language: Option<&str>, + gender: Option<&str>, + limit: usize, + ) -> Result> { + let voices = self.list_voices(language).await?; + let filtered: Vec = match gender { + Some(g) => { + let g_lower = g.to_lowercase(); + voices + .into_iter() + .filter(|v| { + v.gender + .as_deref() + .map(|x| x.to_lowercase() == g_lower) + .unwrap_or(false) + }) + .collect() + } + None => voices, + }; + Ok(filtered.into_iter().take(limit).collect()) + } + + /// 检查是否支持指定能力 + fn supports_capability(&self, cap: Capabilities) -> bool { + self.capabilities().contains(&cap) + } + + /// 检查是否支持指定模型类型 + fn supports_model_type(&self, model_type: ModelType) -> bool { + match model_type { + ModelType::Chat => self.supports_capability(Capabilities::Chat), + ModelType::Image => self.supports_capability(Capabilities::ImageGenerate), + ModelType::Video => self.supports_capability(Capabilities::VideoGenerate), + ModelType::Audio => { + self.supports_capability(Capabilities::AudioTranscribe) + || self.supports_capability(Capabilities::AudioSpeech) + } + } + } +} + +/// 构造"不支持的能力"错误 +fn unsupported(provider: &str, cap: Capabilities) -> AibridgeError { + AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {provider})", cap.as_str()), + } +} + +/// 从 `ProviderConfig` 构造适配器的工厂函数 trait +/// +/// 具体适配器的构造逻辑各异,工厂返回 `Box`。 +/// 对应 Python v1 `AdapterFactory.create`。 +pub type AdapterConstructor = fn(ProviderConfig) -> Result>; + +#[cfg(test)] +mod tests { + use super::*; + use std::collections::HashMap; + + /// 测试用的空适配器(仅实现必需方法,其余走默认实现) + #[allow(dead_code)] + struct DummyAdapter { + config: ProviderConfig, + started: bool, + caps: CapabilitySet, + } + + #[async_trait] + impl Adapter for DummyAdapter { + fn provider_type(&self) -> &str { + "dummy" + } + fn provider_name(&self) -> &str { + "Dummy" + } + fn capabilities(&self) -> CapabilitySet { + self.caps.clone() + } + async fn start(&mut self) -> Result<()> { + self.started = true; + Ok(()) + } + async fn close(&mut self) -> Result<()> { + self.started = false; + Ok(()) + } + async fn list_models(&self, _filter: Option) -> Result> { + Ok(vec![]) + } + } + + fn dummy_config() -> ProviderConfig { + ProviderConfig { + provider_type: "dummy".into(), + api_key: Some("k".into()), + base_url: None, + poll_url: None, + timeout: 300, + max_retries: 3, + retry_delay: 2.0, + enabled: true, + resource_name: None, + deployment_id: None, + api_version: None, + extra: HashMap::new(), + } + } + + #[tokio::test] + async fn default_chat_returns_unsupported_when_no_capability() { + let mut adapter = DummyAdapter { + config: dummy_config(), + started: false, + caps: CapabilitySet::new(), + }; + adapter.start().await.unwrap(); + let req = ChatRequest::builder("m", vec![]).build(); + let result = adapter.chat(req).await; + assert!(matches!( + result, + Err(AibridgeError::UnsupportedCapability { .. }) + )); + } + + #[tokio::test] + async fn default_speech_returns_unsupported() { + let adapter = DummyAdapter { + config: dummy_config(), + started: true, + caps: CapabilitySet::new(), + }; + let req = SpeechRequest::builder("tts-1", "hi", "alloy").build(); + let result = adapter.speech(req).await; + assert!(matches!( + result, + Err(AibridgeError::UnsupportedCapability { .. }) + )); + } + + #[tokio::test] + async fn default_video_poll_returns_unsupported() { + let adapter = DummyAdapter { + config: dummy_config(), + started: true, + caps: CapabilitySet::new(), + }; + let result = adapter.video_poll("t-1", "m").await; + assert!(matches!( + result, + Err(AibridgeError::UnsupportedCapability { .. }) + )); + } + + #[tokio::test] + async fn supports_capability_checks_set() { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + let adapter = DummyAdapter { + config: dummy_config(), + started: true, + caps, + }; + assert!(adapter.supports_capability(Capabilities::Chat)); + assert!(!adapter.supports_capability(Capabilities::Embedding)); + } + + #[tokio::test] + async fn supports_model_type_maps_correctly() { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ImageGenerate); + let adapter = DummyAdapter { + config: dummy_config(), + started: true, + caps, + }; + assert!(adapter.supports_model_type(ModelType::Image)); + assert!(!adapter.supports_model_type(ModelType::Chat)); + } + + #[tokio::test] + async fn recommend_voices_filters_by_gender() { + #[allow(dead_code)] + struct VoiceAdapter { + config: ProviderConfig, + caps: CapabilitySet, + } + #[async_trait] + impl Adapter for VoiceAdapter { + fn provider_type(&self) -> &str { + "voice" + } + fn provider_name(&self) -> &str { + "Voice" + } + fn capabilities(&self) -> CapabilitySet { + self.caps.clone() + } + async fn start(&mut self) -> Result<()> { + Ok(()) + } + async fn close(&mut self) -> Result<()> { + Ok(()) + } + async fn list_models(&self, _: Option) -> Result> { + Ok(vec![]) + } + async fn list_voices(&self, _lang: Option<&str>) -> Result> { + Ok(vec![ + VoiceInfo::builder() + .short_name("v1") + .gender("Female") + .build(), + VoiceInfo::builder().short_name("v2").gender("Male").build(), + ]) + } + } + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ListVoices); + let adapter = VoiceAdapter { + config: dummy_config(), + caps, + }; + let female = adapter + .recommend_voices(None, Some("female"), 10) + .await + .unwrap(); + assert_eq!(female.len(), 1); + assert_eq!(female[0].short_name.as_deref(), Some("v1")); + + let limited = adapter.recommend_voices(None, None, 1).await.unwrap(); + assert_eq!(limited.len(), 1); + } + + #[test] + fn capabilities_as_str_matches_python_constants() { + assert_eq!(Capabilities::Chat.as_str(), "chat"); + assert_eq!(Capabilities::ChatStream.as_str(), "chat_stream"); + assert_eq!(Capabilities::ImageGenerate.as_str(), "image"); + assert_eq!(Capabilities::VideoGenerate.as_str(), "video"); + assert_eq!(Capabilities::Embedding.as_str(), "embedding"); + assert_eq!(Capabilities::AudioTranscribe.as_str(), "audio_transcribe"); + assert_eq!(Capabilities::AudioSpeech.as_str(), "audio_speech"); + assert_eq!(Capabilities::ListVoices.as_str(), "list_voices"); + } + + #[tokio::test] + async fn default_requires_api_key_true() { + let adapter = DummyAdapter { + config: dummy_config(), + started: true, + caps: CapabilitySet::new(), + }; + assert!(adapter.requires_api_key()); + } +} diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs new file mode 100644 index 0000000..100cfc5 --- /dev/null +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -0,0 +1,170 @@ +//! 适配器工厂 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/factory.py`。 +//! +//! 设计要点(与设计文档 5.2 节一致): +//! - 用编译期显式 `match` 替代 Python 运行时注册(更静态、更安全) +//! - 新增适配器=加 match 分支 +//! - 阶段 0.4 暂只占位分支(返 ProviderNotFound),具体适配器阶段 1 起填充 + +use crate::adapter::Adapter; +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; + +/// 已支持的 provider 列表(用于错误信息与测试) +/// +/// 阶段 1 起逐步填充实际支持的 provider 名。 +pub const KNOWN_PROVIDERS: &[&str] = &[ + "openai", + "agnes", + "volcengine_cv", + "gemini", + // 阶段 2 起补充: + "azure", + "anthropic", + "runway", + "pika", + "kling", + "stability", + "chinese", + "aggregation_platforms", + "edge-tts", + "elevenlabs", + "cartesia", + "deepgram", + "assemblyai", +]; + +/// 根据配置创建适配器实例 +/// +/// 对应 Python v1 `AdapterFactory.create`。 +/// 阶段 0.4:所有分支均为占位(返 ProviderNotFound),具体适配器在阶段 1 起填充。 +pub fn create_adapter(config: ProviderConfig) -> Result> { + let provider = config.provider_type.as_str(); + match provider { + // 阶段 1 MVP 适配器(阶段 1.0 起填充实际构造逻辑) + "openai" | "agnes" | "volcengine_cv" | "gemini" => { + // TODO(阶段 1): 引入 adapters::openai::OpenAiAdapter 等具体实现 + Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 1 待实现)"), + }) + } + // 阶段 2 适配器占位 + "azure" + | "anthropic" + | "runway" + | "pika" + | "kling" + | "stability" + | "chinese" + | "aggregation_platforms" + | "edge-tts" + | "elevenlabs" + | "cartesia" + | "deepgram" + | "assemblyai" => Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 2 待实现)"), + }), + // 未知 provider + _ => Err(AibridgeError::provider_not_found(format!( + "{provider}(未知 provider,支持:{})", + KNOWN_PROVIDERS.join(", ") + ))), + } +} + +/// 检查 provider 是否已被工厂识别(不一定已实现) +pub fn is_known_provider(provider: &str) -> bool { + KNOWN_PROVIDERS.contains(&provider) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + + fn config_for(provider: &str) -> ProviderConfig { + ProviderConfig::from_options(provider, ClientOptions::builder().api_key("k").build()) + } + + #[test] + fn create_unknown_provider_returns_error() { + let result = create_adapter(config_for("nonexistent")); + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn create_openai_returns_pending_for_phase0() { + let result = create_adapter(config_for("openai")); + // 阶段 0.4:占位返 ProviderNotFound(待阶段 1 实现) + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + if let Err(AibridgeError::ProviderNotFound { provider }) = result { + assert!(provider.contains("阶段 1")); + } + } + + #[test] + fn create_agnes_returns_pending() { + let result = create_adapter(config_for("agnes")); + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn create_volcengine_cv_returns_pending() { + let result = create_adapter(config_for("volcengine_cv")); + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn create_gemini_returns_pending() { + let result = create_adapter(config_for("gemini")); + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn create_phase2_adapter_returns_phase2_message() { + let result = create_adapter(config_for("anthropic")); + if let Err(AibridgeError::ProviderNotFound { provider }) = result { + assert!(provider.contains("阶段 2")); + } else { + panic!("应为 ProviderNotFound"); + } + } + + #[test] + fn is_known_provider_recognizes_known() { + assert!(is_known_provider("openai")); + assert!(is_known_provider("edge-tts")); + assert!(is_known_provider("assemblyai")); + } + + #[test] + fn is_known_provider_rejects_unknown() { + assert!(!is_known_provider("nonexistent")); + } + + #[test] + fn error_for_unknown_mentions_known_providers() { + let result = create_adapter(config_for("xxx")); + if let Err(AibridgeError::ProviderNotFound { provider }) = result { + assert!(provider.contains("openai")); + } else { + panic!("应为 ProviderNotFound"); + } + } +} diff --git a/crates/aibridge-core/src/adapter/mod.rs b/crates/aibridge-core/src/adapter/mod.rs new file mode 100644 index 0000000..7edc7e3 --- /dev/null +++ b/crates/aibridge-core/src/adapter/mod.rs @@ -0,0 +1,15 @@ +//! 适配器层 +//! +//! 定义 `Adapter` trait、`Capabilities` 枚举与适配器工厂。 +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/base.py` + `agn/adapters/factory.py`。 +//! +//! 设计要点(与设计文档 5.2 节一致): +//! - `Adapter` trait 用 `#[async_trait]`,不支持的方法默认返 `UnsupportedCapability` +//! - 工厂用编译期显式 `match`(替代 Python 运行时注册),新增适配器=加 match 分支 +//! - 具体适配器在阶段 1 起填充 `adapters/` + +pub mod base; +pub mod factory; + +pub use base::{Adapter, Capabilities, CapabilitySet, ChatStream}; +pub use factory::create_adapter; diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs new file mode 100644 index 0000000..207295f --- /dev/null +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -0,0 +1,11 @@ +//! 具体 Provider 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/*.py`(14 个适配器)。 +//! +//! **阶段 0.4**:本模块仅占位,具体适配器在阶段 1 起填充: +//! - 阶段 1(MVP):openai / agnes / volcengine_cv / gemini +//! - 阶段 2a(兼容族):azure / aggregation_platforms / additional_models / more_models / emerging_models / chinese +//! - 阶段 2b(独立协议):anthropic / stability / runway / pika / kling +//! - 阶段 2c(音频):edge-tts / elevenlabs / cartesia / deepgram / assemblyai +//! +//! 各适配器实现 `adapter::Adapter` trait,由 `adapter::create_adapter` 工厂分发。 diff --git a/crates/aibridge-core/src/client.rs b/crates/aibridge-core/src/client.rs new file mode 100644 index 0000000..0e95d27 --- /dev/null +++ b/crates/aibridge-core/src/client.rs @@ -0,0 +1,194 @@ +//! 统一客户端 +//! +//! 提供统一的 API 接口,是用户使用 SDK 的唯一入口。 +//! 对应 Python v1 (agn-sdk) 的 `agn/client.py`。 +//! +//! 设计要点(与设计文档 5.3 节一致): +//! - `Client` 持有 `Box`,方法签名改显式 Request struct +//! - 可选参数通过 Request Builder(替代 `**kwargs`) +//! - 不保留 Python 的 `Options` 中间层 + +use crate::adapter::{create_adapter, Adapter, ChatStream}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; +use crate::model::{ + ChatCompletion, ChatRequest, EmbedRequest, EmbeddingResult, ImageRequest, ImageResult, + SpeechRequest, SpeechResult, TranscribeRequest, TranscriptionResult, VideoRequest, VideoStatus, + VideoTask, +}; + +/// 统一客户端 +/// +/// 对应 Python v1 `Client`。是用户使用 SDK 的唯一入口。 +/// +/// # 示例 +/// ```ignore +/// let client = Client::new("openai", ClientOptions::builder().api_key("sk-xxx").build())?; +/// client.start().await?; +/// let resp = client.chat( +/// ChatRequest::builder("gpt-4o", vec![ChatMessage::user("Hello!")]).build() +/// ).await?; +/// ``` +pub struct Client { + provider_type: String, + adapter: Box, +} + +impl Client { + /// 创建客户端 + /// + /// `provider` 为 Provider 类型(如 "agnes"、"openai")。 + /// `opts` 为连接选项(api_key / base_url / timeout 等)。 + /// + /// 会自动合并环境变量(`AIBRIDGE_{PROVIDER}_*` / `AGN_{PROVIDER}_*`)。 + pub fn new(provider: &str, opts: ClientOptions) -> Result { + let opts = opts.merge_env(provider); + let config = ProviderConfig::from_options(provider, opts); + + // 校验:requires_api_key 时必须有 api_key + // 阶段 0.4 无法预知适配器的 requires_api_key(适配器未实现), + // 暂按"有则需 key"的保守策略:未知 provider 也要求 key + config.validate(true)?; + + let adapter = create_adapter(config.clone()).map_err(|e| match e { + AibridgeError::ProviderNotFound { provider: p } => { + AibridgeError::ProviderNotFound { provider: p } + } + other => other, + })?; + + Ok(Self { + provider_type: provider.to_string(), + adapter, + }) + } + + /// 启动客户端(初始化适配器) + pub async fn start(&mut self) -> Result<()> { + self.adapter.start().await + } + + /// 关闭客户端(释放资源) + pub async fn close(&mut self) -> Result<()> { + self.adapter.close().await + } + + /// Provider 类型 + pub fn provider_type(&self) -> &str { + &self.provider_type + } + + /// 文本对话 + pub async fn chat(&self, req: ChatRequest) -> Result { + self.adapter.chat(req).await + } + + /// 流式文本对话 + pub async fn chat_stream(&self, req: ChatRequest) -> Result { + self.adapter.chat_stream(req).await + } + + /// 图像生成 + pub async fn image_generate(&self, req: ImageRequest) -> Result { + self.adapter.image_generate(req).await + } + + /// 创建视频生成任务 + pub async fn video_create(&self, req: VideoRequest) -> Result { + self.adapter.video_create(req).await + } + + /// 查询视频任务状态 + pub async fn video_poll(&self, task_id: &str, model: &str) -> Result { + self.adapter.video_poll(task_id, model).await + } + + /// 文本嵌入 + pub async fn embed(&self, req: EmbedRequest) -> Result { + self.adapter.embed(req).await + } + + /// 语音转文字 + pub async fn transcribe(&self, req: TranscribeRequest) -> Result { + self.adapter.transcribe(req).await + } + + /// 文字转语音 + pub async fn speech(&self, req: SpeechRequest) -> Result { + self.adapter.speech(req).await + } + + /// 获取可用模型列表 + pub async fn list_models(&self, filter: Option) -> Result> { + self.adapter.list_models(filter).await + } + + /// 获取 Provider 可用音色列表 + pub async fn list_voices(&self, language: Option<&str>) -> Result> { + self.adapter.list_voices(language).await + } + + /// 推荐可用音色(按语言/性别过滤) + pub async fn recommend_voices( + &self, + language: Option<&str>, + gender: Option<&str>, + limit: usize, + ) -> Result> { + self.adapter.recommend_voices(language, gender, limit).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use std::env; + use std::sync::Mutex; + + /// 串行化 env 测试(env 变量是进程级共享,并行执行会互相污染) + static ENV_LOCK: Mutex<()> = Mutex::new(()); + + #[test] + fn new_rejects_missing_api_key() { + let result = Client::new("openai", ClientOptions::default()); + // 缺 api_key + 阶段 0 占位:ValidationError 优先(validate 在 create_adapter 前) + assert!(result.is_err()); + } + + #[test] + fn new_with_key_reaches_factory_and_returns_provider_not_found() { + // 阶段 0.4:工厂占位返 ProviderNotFound + let result = Client::new("openai", ClientOptions::builder().api_key("sk-xxx").build()); + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn new_unknown_provider_returns_provider_not_found() { + let result = Client::new("nonexistent", ClientOptions::builder().api_key("k").build()); + // validate(true) 通过(key 存在),然后 create_adapter 返 ProviderNotFound + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + } + + #[test] + fn new_with_env_var_api_key() { + let _guard = ENV_LOCK.lock().unwrap(); + env::set_var("AIBRIDGE_ENVTEST_API_KEY", "env-key"); + let result = Client::new( + "envtest", + ClientOptions::default(), // 无 api_key,从环境变量补 + ); + // envtest 不是已知 provider,validate 通过后工厂返 ProviderNotFound + assert!(matches!( + result, + Err(AibridgeError::ProviderNotFound { .. }) + )); + env::remove_var("AIBRIDGE_ENVTEST_API_KEY"); + } +} diff --git a/crates/aibridge-core/src/config.rs b/crates/aibridge-core/src/config.rs new file mode 100644 index 0000000..44635c1 --- /dev/null +++ b/crates/aibridge-core/src/config.rs @@ -0,0 +1,489 @@ +//! 配置管理 +//! +//! 定义 `ClientOptions` 与 `ProviderConfig`,并支持从环境变量加载配置。 +//! 对应 Python v1 (agn-sdk) 的 `agn/core/config.py`。 +//! +//! 环境变量兼容性:保留对 `AGN_API_KEY` / `AGN_BASE_URL` 等老前缀的读取, +//! 同时新增 `AIBRIDGE_API_KEY` / `AIBRIDGE_BASE_URL` 前缀,迁移期两者并存。 + +use std::collections::HashMap; +use std::env; +use std::time::Duration; + +use serde::{Deserialize, Serialize}; + +use crate::error::{AibridgeError, Result}; + +/// 客户端全局选项 +/// +/// 用于配置单个 Client 的连接与重试参数。 +/// 对应 Python v1 `Config` 类与 `Client.__init__` 的连接参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ClientOptions { + /// API Key(免费 Provider 如 Edge TTS 可为 None) + #[serde(default)] + pub api_key: Option, + + /// API Base URL(可选,部分 Provider 有默认值) + #[serde(default)] + pub base_url: Option, + + /// 轮询 URL(视频生成任务状态用,部分 Provider 需要) + #[serde(default)] + pub poll_url: Option, + + /// 请求超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout: u64, + + /// 最大重试次数 + #[serde(default = "default_max_retries")] + pub max_retries: u32, + + /// 重试初始延迟(秒) + #[serde(default = "default_retry_delay")] + pub retry_delay: f64, + + /// 默认 Provider(Router 路由失败时兜底) + #[serde(default)] + pub default_provider: Option, + + /// 额外配置(厂商特定配置,如 embed_model / speech_model 等) + #[serde(default)] + pub extra: HashMap, +} + +impl Default for ClientOptions { + fn default() -> Self { + Self { + api_key: None, + base_url: None, + poll_url: None, + timeout: default_timeout(), + max_retries: default_max_retries(), + retry_delay: default_retry_delay(), + default_provider: None, + extra: HashMap::new(), + } + } +} + +impl ClientOptions { + /// 创建一个新的 Builder + pub fn builder() -> ClientOptionsBuilder { + ClientOptionsBuilder::default() + } + + /// 从环境变量加载配置,再与 `self` 合并(`self` 优先级高于环境变量) + /// + /// 读取顺序:`AIBRIDGE_{PROVIDER}_*` > `AGN_{PROVIDER}_*` > 通用 `AIBRIDGE_API_KEY` / `AGN_API_KEY`。 + /// `provider` 用于构造 provider 专属环境变量名(大写)。 + pub fn merge_env(mut self, provider: &str) -> Self { + let upper = provider.to_uppercase(); + + // api_key:provider 专属优先,再回退到通用 + if self.api_key.is_none() { + self.api_key = env_var(&format!("AIBRIDGE_{upper}_API_KEY")) + .or_else(|| env_var(&format!("AGN_{upper}_API_KEY"))) + .or_else(|| env_var("AIBRIDGE_API_KEY")) + .or_else(|| env_var("AGN_API_KEY")); + } + + // base_url + if self.base_url.is_none() { + self.base_url = env_var(&format!("AIBRIDGE_{upper}_BASE_URL")) + .or_else(|| env_var(&format!("AGN_{upper}_BASE_URL"))) + .or_else(|| env_var("AIBRIDGE_BASE_URL")) + .or_else(|| env_var("AGN_BASE_URL")); + } + + // poll_url(视频轮询地址,仅部分 Provider 用) + if self.poll_url.is_none() { + self.poll_url = env_var(&format!("AIBRIDGE_{upper}_POLL_URL")) + .or_else(|| env_var(&format!("AGN_{upper}_POLL_URL"))); + } + + self + } + + /// 获取超时时间对应的 `Duration` + pub fn timeout_duration(&self) -> Duration { + Duration::from_secs(self.timeout) + } +} + +/// `ClientOptions` 的 Builder +/// +/// 用法: +/// ```ignore +/// let opts = ClientOptions::builder() +/// .api_key("sk-xxx") +/// .base_url("https://api.example.com") +/// .timeout(120) +/// .build(); +/// ``` +#[derive(Debug, Default, Clone)] +pub struct ClientOptionsBuilder { + inner: ClientOptions, +} + +impl ClientOptionsBuilder { + pub fn api_key(mut self, api_key: impl Into) -> Self { + self.inner.api_key = Some(api_key.into()); + self + } + + pub fn base_url(mut self, base_url: impl Into) -> Self { + self.inner.base_url = Some(base_url.into()); + self + } + + pub fn poll_url(mut self, poll_url: impl Into) -> Self { + self.inner.poll_url = Some(poll_url.into()); + self + } + + pub fn timeout(mut self, timeout: u64) -> Self { + self.inner.timeout = timeout; + self + } + + pub fn max_retries(mut self, max_retries: u32) -> Self { + self.inner.max_retries = max_retries; + self + } + + pub fn retry_delay(mut self, retry_delay: f64) -> Self { + self.inner.retry_delay = retry_delay; + self + } + + pub fn default_provider(mut self, provider: impl Into) -> Self { + self.inner.default_provider = Some(provider.into()); + self + } + + /// 插入一个额外配置项(厂商特定参数) + pub fn extra(mut self, key: impl Into, value: impl Into) -> Self { + self.inner.extra.insert(key.into(), value.into()); + self + } + + pub fn build(self) -> ClientOptions { + self.inner + } +} + +/// Provider 配置 +/// +/// 描述单个 AI 模型提供商的连接参数。 +/// 对应 Python v1 `agn/models/common.py` 的 `ProviderConfig`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderConfig { + /// Provider 类型标识,如 "agnes"、"openai" + pub provider_type: String, + + /// API Key(免费 Provider 可为 None) + #[serde(default)] + pub api_key: Option, + + /// API Base URL(可选,部分 Provider 有默认值) + #[serde(default)] + pub base_url: Option, + + /// 轮询 URL(视频生成任务状态用) + #[serde(default)] + pub poll_url: Option, + + /// 请求超时时间(秒) + #[serde(default = "default_timeout")] + pub timeout: u64, + + /// 最大重试次数 + #[serde(default = "default_max_retries")] + pub max_retries: u32, + + /// 重试延迟(秒) + #[serde(default = "default_retry_delay")] + pub retry_delay: f64, + + /// 是否启用该 Provider(Router 用) + #[serde(default = "default_enabled")] + pub enabled: bool, + + /// Azure 专用字段:资源名称 + #[serde(default)] + pub resource_name: Option, + + /// Azure 专用字段:部署 ID + #[serde(default)] + pub deployment_id: Option, + + /// API 版本(Azure 等用) + #[serde(default)] + pub api_version: Option, + + /// 额外配置(厂商特定配置) + #[serde(default)] + pub extra: HashMap, +} + +impl ProviderConfig { + /// 从 `ClientOptions` 与 provider 类型构造配置 + /// + /// `opts` 的字段优先;opts.api_key 为 None 时保留 None(由调用方按 + /// `requires_api_key` 决定是否报错)。 + pub fn from_options(provider_type: impl Into, opts: ClientOptions) -> Self { + Self { + provider_type: provider_type.into(), + api_key: opts.api_key, + base_url: opts.base_url, + poll_url: opts.poll_url, + timeout: opts.timeout, + max_retries: opts.max_retries, + retry_delay: opts.retry_delay, + enabled: true, + resource_name: opts + .extra + .get("resource_name") + .and_then(|v| v.as_str()) + .map(str::to_owned), + deployment_id: opts + .extra + .get("deployment_id") + .and_then(|v| v.as_str()) + .map(str::to_owned), + api_version: opts + .extra + .get("api_version") + .and_then(|v| v.as_str()) + .map(str::to_owned), + extra: opts.extra, + } + } + + /// 校验配置:必填字段检查等 + /// + /// `requires_api_key` 为 true 时,api_key 必须非空。 + pub fn validate(&self, requires_api_key: bool) -> Result<()> { + if self.provider_type.trim().is_empty() { + return Err(AibridgeError::validation("provider_type 不能为空")); + } + if requires_api_key { + match &self.api_key { + Some(k) if !k.trim().is_empty() => Ok(()), + _ => Err(AibridgeError::validation(format!( + "Provider '{}' 需要 API key", + self.provider_type + ))), + } + } else { + Ok(()) + } + } +} + +/// 读取环境变量,返回非空字符串(空串视为未设置) +fn env_var(key: &str) -> Option { + env::var(key) + .ok() + .map(|s| s.trim().to_owned()) + .filter(|s| !s.is_empty()) +} + +fn default_timeout() -> u64 { + 300 +} + +fn default_max_retries() -> u32 { + 3 +} + +fn default_retry_delay() -> f64 { + 2.0 +} + +fn default_enabled() -> bool { + true +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::Mutex; + + /// 串行化所有 env 测试(env 变量是进程级共享,并行执行会互相污染) + static ENV_LOCK: Mutex<()> = Mutex::new(()); + + #[test] + fn default_options() { + let opts = ClientOptions::default(); + assert_eq!(opts.timeout, 300); + assert_eq!(opts.max_retries, 3); + assert!((opts.retry_delay - 2.0).abs() < f64::EPSILON); + assert!(opts.api_key.is_none()); + } + + #[test] + fn builder_sets_fields() { + let opts = ClientOptions::builder() + .api_key("sk-xxx") + .base_url("https://api.example.com") + .timeout(120) + .max_retries(5) + .retry_delay(1.0) + .default_provider("openai") + .extra("k", "v") + .build(); + assert_eq!(opts.api_key.as_deref(), Some("sk-xxx")); + assert_eq!(opts.base_url.as_deref(), Some("https://api.example.com")); + assert_eq!(opts.timeout, 120); + assert_eq!(opts.max_retries, 5); + assert!((opts.retry_delay - 1.0).abs() < f64::EPSILON); + assert_eq!(opts.default_provider.as_deref(), Some("openai")); + assert_eq!(opts.extra.get("k").and_then(|v| v.as_str()), Some("v")); + } + + #[test] + fn timeout_duration() { + let opts = ClientOptions::builder().timeout(42).build(); + assert_eq!(opts.timeout_duration(), Duration::from_secs(42)); + } + + #[test] + fn merge_env_uses_existing_first() { + let _guard = ENV_LOCK.lock().unwrap(); + env::set_var("AIBRIDGE_TEST_MERGE_API_KEY", "env-key"); + env::set_var("AIBRIDGE_API_KEY", "global-key"); + let opts = ClientOptions::builder() + .api_key("explicit-key") + .build() + .merge_env("test_merge"); + assert_eq!(opts.api_key.as_deref(), Some("explicit-key")); + env::remove_var("AIBRIDGE_TEST_MERGE_API_KEY"); + env::remove_var("AIBRIDGE_API_KEY"); + } + + #[test] + fn merge_env_falls_back_to_provider_specific() { + let _guard = ENV_LOCK.lock().unwrap(); + env::set_var("AIBRIDGE_FALLBACK_API_KEY", "provider-key"); + let opts = ClientOptions::default().merge_env("fallback"); + assert_eq!(opts.api_key.as_deref(), Some("provider-key")); + env::remove_var("AIBRIDGE_FALLBACK_API_KEY"); + } + + #[test] + fn merge_env_falls_back_to_global_agn_prefix() { + let _guard = ENV_LOCK.lock().unwrap(); + env::set_var("AGN_API_KEY", "agn-global"); + let opts = ClientOptions::default().merge_env("nobody_has_this"); + assert_eq!(opts.api_key.as_deref(), Some("agn-global")); + env::remove_var("AGN_API_KEY"); + } + + #[test] + fn merge_env_base_url_provider_specific() { + let _guard = ENV_LOCK.lock().unwrap(); + env::set_var("AIBRIDGE_URLPROV_BASE_URL", "https://provider.example.com"); + let opts = ClientOptions::default().merge_env("urlprov"); + assert_eq!( + opts.base_url.as_deref(), + Some("https://provider.example.com") + ); + env::remove_var("AIBRIDGE_URLPROV_BASE_URL"); + } + + #[test] + fn merge_env_empty_string_treated_as_unset() { + let _guard = ENV_LOCK.lock().unwrap(); + // 确保全局回退变量未设置(隔离其他测试的污染) + env::remove_var("AIBRIDGE_API_KEY"); + env::remove_var("AGN_API_KEY"); + env::set_var("AIBRIDGE_EMPTY_API_KEY", ""); + let opts = ClientOptions::default().merge_env("empty"); + assert!(opts.api_key.is_none()); + env::remove_var("AIBRIDGE_EMPTY_API_KEY"); + } + + #[test] + fn provider_config_from_options() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("u") + .timeout(99) + .build(); + let cfg = ProviderConfig::from_options("openai", opts); + assert_eq!(cfg.provider_type, "openai"); + assert_eq!(cfg.api_key.as_deref(), Some("k")); + assert_eq!(cfg.base_url.as_deref(), Some("u")); + assert_eq!(cfg.timeout, 99); + assert!(cfg.enabled); + } + + #[test] + fn provider_config_validate_requires_key_when_needed() { + let cfg = ProviderConfig { + provider_type: "openai".into(), + api_key: None, + ..Default::default() + }; + assert!(cfg.validate(true).is_err()); + + let cfg2 = ProviderConfig { + provider_type: "openai".into(), + api_key: Some("k".into()), + ..Default::default() + }; + assert!(cfg2.validate(true).is_ok()); + } + + #[test] + fn provider_config_validate_skips_key_when_not_required() { + let cfg = ProviderConfig { + provider_type: "edge-tts".into(), + api_key: None, + ..Default::default() + }; + assert!(cfg.validate(false).is_ok()); + } + + #[test] + fn provider_config_validate_rejects_empty_provider_type() { + let cfg = ProviderConfig { + provider_type: " ".into(), + api_key: Some("k".into()), + ..Default::default() + }; + assert!(cfg.validate(false).is_err()); + } + + impl Default for ProviderConfig { + fn default() -> Self { + Self { + provider_type: String::new(), + api_key: None, + base_url: None, + poll_url: None, + timeout: default_timeout(), + max_retries: default_max_retries(), + retry_delay: default_retry_delay(), + enabled: true, + resource_name: None, + deployment_id: None, + api_version: None, + extra: HashMap::new(), + } + } + } + + #[test] + fn provider_config_serialize_roundtrip() { + let cfg = + ProviderConfig::from_options("agnes", ClientOptions::builder().api_key("k").build()); + let json = serde_json::to_string(&cfg).unwrap(); + let back: ProviderConfig = serde_json::from_str(&json).unwrap(); + assert_eq!(back.provider_type, "agnes"); + assert_eq!(back.api_key.as_deref(), Some("k")); + } +} diff --git a/crates/aibridge-core/src/error.rs b/crates/aibridge-core/src/error.rs new file mode 100644 index 0000000..30072b9 --- /dev/null +++ b/crates/aibridge-core/src/error.rs @@ -0,0 +1,409 @@ +//! 错误类型定义 +//! +//! 定义统一的错误枚举 `AibridgeError`,用于映射各 Provider 的特定错误。 +//! 对应 Python v1 (agn-sdk) 的 `agn/core/errors.py`,分类保持一致。 +//! +//! 设计文档 9.1 节规定的错误类别: +//! Authentication / RateLimit / Validation / ModelNotFound / Api / Network / Timeout +//! / UnsupportedCapability / ProviderNotFound。 +//! +//! 迁移对照:Python v1 的 `AGNError` → Rust 的 `AibridgeError`(子类名不变)。 + +use std::time::Duration; + +/// SDK 统一错误类型 +/// +/// 所有 SDK 操作返回的错误枚举,按错误性质分类。 +/// 各子类对应 Python v1 的同名错误类。 +#[derive(Debug, thiserror::Error)] +pub enum AibridgeError { + /// 认证错误 - API Key 无效、过期或没有权限 + /// + /// 对应 Python v1 `AuthenticationError`。 + #[error("认证失败: {message}")] + Authentication { message: String }, + + /// 限流错误 - 请求频率超过限制 + /// + /// 对应 Python v1 `RateLimitError`。`retry_after` 为服务端建议的等待时间(秒)。 + #[error("限流: {message}")] + RateLimit { + message: String, + /// 服务端建议的等待时间(秒),None 表示未提供 + retry_after: Option, + }, + + /// 参数校验错误 - 请求参数不合法 + /// + /// 对应 Python v1 `ValidationError`。`details` 携带结构化错误细节。 + #[error("参数校验错误: {message}")] + Validation { + message: String, + /// 结构化错误细节(字段级错误、约束违反等) + details: serde_json::Value, + }, + + /// 模型不存在错误 - 请求的模型不可用 + /// + /// 对应 Python v1 `ModelNotFoundError`。 + #[error("模型不存在: {model}")] + ModelNotFound { model: String }, + + /// API 调用错误 - 来自 Provider API 的错误响应 + /// + /// 对应 Python v1 `APIError`。`status` 为 HTTP 状态码。 + #[error("API 调用错误: {message}")] + Api { + /// HTTP 状态码 + status: u16, + message: String, + }, + + /// 网络错误 - 网络连接问题 + /// + /// 对应 Python v1 `NetworkError`。由 `reqwest::Error` 自动转换而来。 + #[error("网络错误: {0}")] + Network(#[from] reqwest::Error), + + /// 超时错误 - 请求超时 + /// + /// 对应 Python v1 `TimeoutError`。 + #[error("超时")] + Timeout, + + /// 不支持的能力错误 - Provider 不支持请求的能力 + /// + /// 对应 Python v1 `UnsupportedCapabilityError`。 + /// `capability` 为不支持的能力标识(如 "chat_stream"、"embedding")。 + #[error("不支持的能力: {capability}")] + UnsupportedCapability { capability: String }, + + /// Provider 不存在错误 - 请求的 Provider 类型不可用 + /// + /// 对应 Python v1 `ProviderNotFoundError`。 + #[error("Provider 不存在: {provider}")] + ProviderNotFound { provider: String }, + + /// 音色不可用错误 - 请求的 voice 已下线或不存在 + /// + /// 对应 Python v1 `VoiceNotAvailableError`。重试无意义,应换音色。 + #[error("音色不可用: {message}")] + VoiceNotAvailable { message: String }, + + /// 服务不可用错误 - Provider 服务端临时不可用(限流/抖动),可重试 + /// + /// 对应 Python v1 `ServiceUnavailableError`。 + #[error("服务暂时不可用: {message}")] + ServiceUnavailable { message: String }, +} + +impl AibridgeError { + /// 创建认证错误 + pub fn authentication(message: impl Into) -> Self { + Self::Authentication { + message: message.into(), + } + } + + /// 创建限流错误(无 retry_after) + pub fn rate_limit(message: impl Into) -> Self { + Self::RateLimit { + message: message.into(), + retry_after: None, + } + } + + /// 创建限流错误(带 retry_after) + pub fn rate_limit_with_retry(message: impl Into, retry_after: f64) -> Self { + Self::RateLimit { + message: message.into(), + retry_after: Some(retry_after), + } + } + + /// 创建参数校验错误 + pub fn validation(message: impl Into) -> Self { + Self::Validation { + message: message.into(), + details: serde_json::Value::Null, + } + } + + /// 创建带结构化细节的参数校验错误 + pub fn validation_with_details(message: impl Into, details: serde_json::Value) -> Self { + Self::Validation { + message: message.into(), + details, + } + } + + /// 创建模型不存在错误 + pub fn model_not_found(model: impl Into) -> Self { + Self::ModelNotFound { + model: model.into(), + } + } + + /// 创建 API 调用错误 + pub fn api(status: u16, message: impl Into) -> Self { + Self::Api { + status, + message: message.into(), + } + } + + /// 创建不支持的能力错误 + pub fn unsupported_capability(capability: impl Into) -> Self { + Self::UnsupportedCapability { + capability: capability.into(), + } + } + + /// 创建 Provider 不存在错误 + pub fn provider_not_found(provider: impl Into) -> Self { + Self::ProviderNotFound { + provider: provider.into(), + } + } + + /// 创建音色不可用错误 + pub fn voice_not_available(message: impl Into) -> Self { + Self::VoiceNotAvailable { + message: message.into(), + } + } + + /// 创建服务不可用错误 + pub fn service_unavailable(message: impl Into) -> Self { + Self::ServiceUnavailable { + message: message.into(), + } + } + + /// 将 HTTP 状态码映射到对应的错误类型 + /// + /// 对应 Python v1 `map_http_status_to_error`。 + /// + /// 映射规则: + /// - 401/403 → Authentication + /// - 429 → RateLimit + /// - 400 → Validation + /// - 404 → ModelNotFound + /// - 4xx(其他)→ Api + /// - 5xx → Api + pub fn from_http_status(status: u16, body: &str) -> Self { + let details = serde_json::json!({ + "status_code": status, + "response": body, + }); + match status { + 401 | 403 => Self::Authentication { + message: "Authentication failed. Please check your API key.".into(), + }, + 429 => Self::RateLimit { + message: "Rate limit exceeded. Please slow down your requests.".into(), + retry_after: None, + }, + 400 => Self::Validation { + message: "Invalid request. Please check your parameters.".into(), + details, + }, + 404 => Self::ModelNotFound { + model: "The requested model was not found.".into(), + }, + s if (400..500).contains(&s) => Self::Api { + status, + message: format!("Client error: {status}"), + }, + s if (500..600).contains(&s) => Self::Api { + status, + message: format!("Server error: {status}. Please try again later."), + }, + _ => Self::Api { + status, + message: format!("Unexpected status code: {status}"), + }, + } + } + + /// 判断该错误是否可重试 + /// + /// 可重试的错误:RateLimit / Network / Timeout / ServiceUnavailable / Api(5xx)。 + /// 不可重试的错误:Authentication / Validation / ModelNotFound / + /// UnsupportedCapability / ProviderNotFound / VoiceNotAvailable / Api(4xx)。 + pub fn is_retryable(&self) -> bool { + match self { + Self::RateLimit { .. } => true, + Self::Network(_) => true, + Self::Timeout => true, + Self::ServiceUnavailable { .. } => true, + // 5xx 服务端错误可重试 + Self::Api { status, .. } => *status >= 500, + // 其余错误重试无意义 + Self::Authentication { .. } + | Self::Validation { .. } + | Self::ModelNotFound { .. } + | Self::UnsupportedCapability { .. } + | Self::ProviderNotFound { .. } + | Self::VoiceNotAvailable { .. } => false, + } + } + + /// 错误的稳定标识码(snake_case),用于序列化到 FFI 的 last_error JSON + /// + /// 对应 Python v1 各错误类的 `code` 字段。 + pub fn code(&self) -> &'static str { + match self { + Self::Authentication { .. } => "authentication_error", + Self::RateLimit { .. } => "rate_limit_error", + Self::Validation { .. } => "validation_error", + Self::ModelNotFound { .. } => "model_not_found", + Self::Api { .. } => "api_error", + Self::Network(_) => "network_error", + Self::Timeout => "timeout_error", + Self::UnsupportedCapability { .. } => "unsupported_capability", + Self::ProviderNotFound { .. } => "provider_not_found", + Self::VoiceNotAvailable { .. } => "voice_not_available", + Self::ServiceUnavailable { .. } => "service_unavailable", + } + } + + /// 若为限流错误,返回 retry_after 对应的等待时长 + pub fn retry_after(&self) -> Option { + match self { + Self::RateLimit { retry_after, .. } => retry_after.map(Duration::from_secs_f64), + _ => None, + } + } +} + +/// SDK Result 类型别名,统一用 `AibridgeError` 作错误 +pub type Result = std::result::Result; + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn from_http_status_401_is_authentication() { + let err = AibridgeError::from_http_status(401, "unauthorized"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn from_http_status_429_is_rate_limit() { + let err = AibridgeError::from_http_status(429, "slow down"); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn from_http_status_400_is_validation() { + let err = AibridgeError::from_http_status(400, "bad request"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn from_http_status_404_is_model_not_found() { + let err = AibridgeError::from_http_status(404, "not found"); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[test] + fn from_http_status_500_is_api_and_retryable() { + let err = AibridgeError::from_http_status(500, "server error"); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + assert!(err.is_retryable()); + } + + #[test] + fn from_http_status_422_is_api_and_not_retryable() { + let err = AibridgeError::from_http_status(422, "unprocessable"); + assert!(matches!(err, AibridgeError::Api { status: 422, .. })); + assert!(!err.is_retryable()); + } + + #[test] + fn rate_limit_is_retryable() { + assert!(AibridgeError::rate_limit("slow").is_retryable()); + assert!(AibridgeError::rate_limit_with_retry("slow", 2.5).is_retryable()); + } + + #[test] + fn timeout_is_retryable() { + assert!(AibridgeError::Timeout.is_retryable()); + } + + #[test] + fn authentication_is_not_retryable() { + assert!(!AibridgeError::authentication("bad key").is_retryable()); + } + + #[test] + fn validation_is_not_retryable() { + assert!(!AibridgeError::validation("bad param").is_retryable()); + } + + #[test] + fn model_not_found_is_not_retryable() { + assert!(!AibridgeError::model_not_found("gpt-x").is_retryable()); + } + + #[test] + fn unsupported_capability_is_not_retryable() { + assert!(!AibridgeError::unsupported_capability("video").is_retryable()); + } + + #[test] + fn provider_not_found_is_not_retryable() { + assert!(!AibridgeError::provider_not_found("foo").is_retryable()); + } + + #[test] + fn voice_not_available_is_not_retryable() { + assert!(!AibridgeError::voice_not_available("offline").is_retryable()); + } + + #[test] + fn service_unavailable_is_retryable() { + assert!(AibridgeError::service_unavailable("temp").is_retryable()); + } + + #[test] + fn retry_after_returns_duration() { + let err = AibridgeError::rate_limit_with_retry("slow", 1.5); + assert_eq!(err.retry_after(), Some(Duration::from_millis(1500))); + } + + #[test] + fn retry_after_none_for_non_rate_limit() { + assert_eq!(AibridgeError::Timeout.retry_after(), None); + } + + #[test] + fn code_is_stable() { + assert_eq!(AibridgeError::Timeout.code(), "timeout_error"); + assert_eq!(AibridgeError::rate_limit("").code(), "rate_limit_error"); + assert_eq!( + AibridgeError::authentication("").code(), + "authentication_error" + ); + } + + #[test] + fn display_includes_message() { + let err = AibridgeError::authentication("bad key"); + assert!(err.to_string().contains("bad key")); + } + + #[test] + fn validation_with_details_carries_json() { + let err = + AibridgeError::validation_with_details("bad", serde_json::json!({"field": "model"})); + match err { + AibridgeError::Validation { details, .. } => { + assert_eq!(details["field"], "model"); + } + _ => panic!("应为 Validation"), + } + } +} diff --git a/crates/aibridge-core/src/http.rs b/crates/aibridge-core/src/http.rs new file mode 100644 index 0000000..a2321b3 --- /dev/null +++ b/crates/aibridge-core/src/http.rs @@ -0,0 +1,284 @@ +//! HTTP 客户端封装 +//! +//! 基于 `reqwest` 封装,提供连接池、统一错误处理等功能。 +//! 对应 Python v1 (agn-sdk) 的 `agn/core/http_client.py`(替代 httpx)。 +//! +//! 设计要点: +//! - 默认启用 HTTP/2(reqwest `http2` feature) +//! - 连接池:`pool_max_idle_per_host` 控制每主机最大空闲连接 +//! - 超时:总超时 + 连接超时 +//! - 统一响应错误映射:`>=400` 通过 `AibridgeError::from_http_status` + +use std::sync::Arc; +use std::time::Duration; + +use reqwest::{Client, Method, Response}; + +use crate::config::ClientOptions; +use crate::error::{AibridgeError, Result}; + +/// 异步 HTTP 客户端 +/// +/// 封装 `reqwest::Client`,提供: +/// - 连接池复用 +/// - 统一错误处理 +/// - 请求/响应日志(tracing) +/// +/// 对应 Python v1 `AsyncHttpClient`。 +#[derive(Clone)] +pub struct HttpClient { + client: Client, + base_url: Option, +} + +impl HttpClient { + /// 构建一个 HTTP 客户端 + /// + /// `opts.base_url` 作为请求 URL 前缀(相对路径会拼接到此)。 + /// `opts.timeout` 为请求总超时;连接超时固定 30 秒。 + /// + /// 注意:reqwest 0.12 的 `ClientBuilder` 不直接支持 base_url, + /// 这里通过 `resolve_url` 在每次请求时手动拼接。 + pub fn new(opts: &ClientOptions) -> Result { + let connect_timeout = Duration::from_secs(30); + let timeout = opts.timeout_duration(); + + let builder = Client::builder() + .timeout(timeout) + .connect_timeout(connect_timeout) + .pool_max_idle_per_host(20) + .https_only(false); + + let client = builder.build().map_err(AibridgeError::from)?; + Ok(Self { + client, + base_url: opts.base_url.clone(), + }) + } + + /// 用显式的 reqwest::Client 构造(测试用) + #[cfg(test)] + pub fn from_client(client: Client, base_url: Option) -> Self { + Self { client, base_url } + } + + /// 返回基础 URL(如有) + pub fn base_url(&self) -> Option<&str> { + self.base_url.as_deref() + } + + /// 解析 URL:相对路径拼接 base_url,绝对路径直接使用 + fn resolve_url(&self, url: &str) -> String { + if url.starts_with("http://") || url.starts_with("https://") { + return url.to_string(); + } + match &self.base_url { + Some(base) => { + let base = base.trim_end_matches('/'); + let path = url.trim_start_matches('/'); + format!("{base}/{path}") + } + None => url.to_string(), + } + } + + /// 发送 GET 请求 + pub async fn get(&self, url: &str) -> Result { + self.request(Method::GET, url).await + } + + /// 发送 POST 请求(JSON body) + pub async fn post_json( + &self, + url: &str, + body: &T, + ) -> Result { + let resp = self + .client + .post(self.resolve_url(url)) + .json(body) + .send() + .await + .map_err(map_reqwest_error)?; + handle_response(resp).await + } + + /// 发送 POST 请求(原始字节 body) + pub async fn post_bytes( + &self, + url: &str, + content_type: &str, + body: bytes::Bytes, + ) -> Result { + let resp = self + .client + .post(self.resolve_url(url)) + .header(reqwest::header::CONTENT_TYPE, content_type) + .body(body) + .send() + .await + .map_err(map_reqwest_error)?; + handle_response(resp).await + } + + /// 发送带自定义请求构造的请求(适配器需要自定义 header/body 时用) + pub async fn request(&self, method: Method, url: &str) -> Result { + let resp = self + .client + .request(method, self.resolve_url(url)) + .send() + .await + .map_err(map_reqwest_error)?; + handle_response(resp).await + } + + /// 发送带 Bearer 认证的 JSON 请求 + pub async fn post_json_authed( + &self, + url: &str, + api_key: &str, + body: &T, + ) -> Result { + let resp = self + .client + .post(self.resolve_url(url)) + .bearer_auth(api_key) + .json(body) + .send() + .await + .map_err(map_reqwest_error)?; + handle_response(resp).await + } + + /// 发送带 Bearer 认证的 GET 请求 + pub async fn get_authed(&self, url: &str, api_key: &str) -> Result { + let resp = self + .client + .get(self.resolve_url(url)) + .bearer_auth(api_key) + .send() + .await + .map_err(map_reqwest_error)?; + handle_response(resp).await + } + + /// 获取底层 reqwest::Client 的引用(适配器需要更细粒度控制时用) + pub fn inner(&self) -> &Client { + &self.client + } +} + +/// 将 reqwest::Error 映射为 AibridgeError +/// +/// 超时 → Timeout;其余 → Network。 +fn map_reqwest_error(err: reqwest::Error) -> AibridgeError { + if err.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(err) + } +} + +/// 处理响应:状态码 >= 400 转为错误 +/// +/// 对应 Python v1 `_handle_response`。 +async fn handle_response(resp: Response) -> Result { + let status = resp.status(); + if status.is_success() { + return Ok(resp); + } + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + Err(AibridgeError::from_http_status(status_code, &body_text)) +} + +/// 共享的 HTTP 客户端句柄 +pub type SharedHttpClient = Arc; + +/// 将 HttpClient 转为共享句柄 +pub fn shared(client: HttpClient) -> SharedHttpClient { + Arc::new(client) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + + fn make_client() -> HttpClient { + HttpClient::new(&ClientOptions::default()).expect("构建客户端失败") + } + + #[test] + fn resolve_url_absolute_passthrough() { + let c = make_client(); + assert_eq!( + c.resolve_url("https://api.example.com/v1"), + "https://api.example.com/v1" + ); + assert_eq!( + c.resolve_url("http://localhost:8080"), + "http://localhost:8080" + ); + } + + #[test] + fn resolve_url_relative_joins_base() { + let opts = ClientOptions::builder() + .base_url("https://api.example.com/") + .build(); + let c = HttpClient::new(&opts).unwrap(); + assert_eq!(c.resolve_url("/v1/chat"), "https://api.example.com/v1/chat"); + assert_eq!(c.resolve_url("v1/chat"), "https://api.example.com/v1/chat"); + } + + #[test] + fn resolve_url_no_base_returns_as_is() { + let c = make_client(); + assert_eq!(c.resolve_url("/v1/chat"), "/v1/chat"); + } + + #[test] + fn base_url_exposed() { + let opts = ClientOptions::builder() + .base_url("https://api.example.com") + .build(); + let c = HttpClient::new(&opts).unwrap(); + assert_eq!(c.base_url(), Some("https://api.example.com")); + } + + #[test] + fn client_constructs_with_pool_and_timeout() { + let opts = ClientOptions::builder().timeout(42).build(); + let c = HttpClient::new(&opts).unwrap(); + // 内部 client 应存在且可用(无法直接断言 timeout,但能拿到 inner 即可) + let _inner = c.inner(); + } + + #[test] + fn client_with_invalid_base_url_still_constructs() { + // base_url 仅作字符串保留(手动拼接),任何字符串都不会导致构造失败 + let opts = ClientOptions::builder() + .base_url("not a url at all") + .build(); + let c = HttpClient::new(&opts); + assert!(c.is_ok()); + if let Ok(c) = c { + assert_eq!(c.base_url(), Some("not a url at all")); + } + } + + #[test] + fn from_http_status_maps_401() { + let err = AibridgeError::from_http_status(401, ""); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn shared_client_wraps_in_arc() { + let c = make_client(); + let s = shared(c); + // Arc 引用计数为 1 + assert_eq!(Arc::strong_count(&s), 1); + } +} diff --git a/crates/aibridge-core/src/lib.rs b/crates/aibridge-core/src/lib.rs index 38a44c3..e6b79b9 100644 --- a/crates/aibridge-core/src/lib.rs +++ b/crates/aibridge-core/src/lib.rs @@ -8,17 +8,17 @@ //! 对应 Python v1 (agn-sdk) 的 agn/ 目录,五层架构保持一致: //! API 层(client) → 路由层(router) → 适配器层(adapter) → 核心层(http/retry/error/config) → 模型层(model) -// 模块声明(阶段 0.2–0.4 逐步填充,暂注释) -// pub mod error; -// pub mod config; -// pub mod http; -// pub mod retry; -// pub mod util; -// pub mod model; -// pub mod adapter; -// pub mod adapters; -// pub mod client; -// pub mod router; +// 模块声明(阶段 0.2–0.4) +pub mod adapter; +pub mod adapters; +pub mod client; +pub mod config; +pub mod error; +pub mod http; +pub mod model; +pub mod retry; +pub mod router; +pub mod util; /// crate 版本号 pub const VERSION: &str = env!("CARGO_PKG_VERSION"); diff --git a/crates/aibridge-core/src/model/audio.rs b/crates/aibridge-core/src/model/audio.rs new file mode 100644 index 0000000..7a03a8e --- /dev/null +++ b/crates/aibridge-core/src/model/audio.rs @@ -0,0 +1,514 @@ +//! 语音数据模型 +//! +//! 定义语音转文字(ASR)与文字转语音(TTS)相关的 serde struct。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/audio.py`。 + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +use crate::model::image::FileInput; + +/// 语音转文字请求 +/// +/// 对应 Python v1 `TranscribeOptions` + 请求参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TranscribeRequest { + /// 模型名称(如 "whisper-1") + pub model: String, + /// 音频文件(路径、URL、base64 或二进制) + pub file: FileInput, + /// 语言代码(如 "zh"、"en") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub language: Option, + /// 提示词(改善专有名词识别、纠正错别字) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub prompt: Option, + /// 响应格式("json" / "text" / "srt" / "vtt" / "verbose_json") + #[serde( + default = "default_transcribe_format", + skip_serializing_if = "is_default_transcribe_format" + )] + pub response_format: String, + /// 温度系数(0-1) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub temperature: Option, + /// 时间戳精度("word" / "segment") + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub timestamp_granularities: Vec, + /// 是否翻译为英文(部分模型支持) + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub translate: bool, + /// 厂商特有参数透传 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +impl TranscribeRequest { + /// 创建 Builder + pub fn builder(model: impl Into, file: FileInput) -> TranscribeRequestBuilder { + TranscribeRequestBuilder { + inner: TranscribeRequest { + model: model.into(), + file, + language: None, + prompt: None, + response_format: default_transcribe_format(), + temperature: None, + timestamp_granularities: Vec::new(), + translate: false, + extra: HashMap::new(), + }, + } + } +} + +/// `TranscribeRequest` 的 Builder +#[derive(Debug, Clone)] +pub struct TranscribeRequestBuilder { + inner: TranscribeRequest, +} + +impl TranscribeRequestBuilder { + pub fn language(mut self, l: impl Into) -> Self { + self.inner.language = Some(l.into()); + self + } + pub fn prompt(mut self, p: impl Into) -> Self { + self.inner.prompt = Some(p.into()); + self + } + pub fn response_format(mut self, r: impl Into) -> Self { + self.inner.response_format = r.into(); + self + } + pub fn temperature(mut self, t: f64) -> Self { + self.inner.temperature = Some(t); + self + } + pub fn timestamp_granularities(mut self, g: Vec) -> Self { + self.inner.timestamp_granularities = g; + self + } + pub fn translate(mut self, t: bool) -> Self { + self.inner.translate = t; + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> TranscribeRequest { + self.inner + } +} + +/// 转写结果 +/// +/// 对应 Python v1 `TranscriptionResult`。 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct TranscriptionResult { + /// 完整转写文本 + pub text: String, + /// 检测到的语言 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub language: Option, + /// 音频时长(秒) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub duration: Option, + /// 分段信息 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub segments: Option>, + /// 词级时间戳 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub words: Option>, + /// 任务类型(transcribe / translate) + #[serde(default = "default_task", skip_serializing_if = "is_default_task")] + pub task: String, + /// 使用统计 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, + /// 使用的模型 ID + #[serde(default, skip_serializing_if = "Option::is_none")] + pub model: Option, +} + +/// 转写分段信息(带时间戳) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TranscriptionSegment { + /// 分段 ID + pub id: u32, + /// 开始时间(秒) + pub start: f64, + /// 结束时间(秒) + pub end: f64, + /// 分段文本 + pub text: String, + /// 分段置信度(0-1) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub confidence: Option, + /// 说话人标识(说话人分离时使用) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub speaker: Option, +} + +/// 转写词级时间戳信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct TranscriptionWord { + /// 词文本 + pub word: String, + /// 开始时间(秒) + pub start: f64, + /// 结束时间(秒) + pub end: f64, + /// 置信度(0-1) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub confidence: Option, +} + +/// 文字转语音请求 +/// +/// 对应 Python v1 `SpeechOptions` + 请求参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct SpeechRequest { + /// 模型名称(如 "tts-1"、"tts-1-hd") + pub model: String, + /// 要合成的文本 + pub input: String, + /// 音色(单个或候选列表用于自动降级) + pub voice: VoiceSpec, + /// 音频输出格式("mp3" / "opus" / "aac" / "flac" / "wav" / "pcm") + #[serde( + default = "default_speech_format", + skip_serializing_if = "is_default_speech_format" + )] + pub response_format: String, + /// 语速(0.25-4.0,默认 1.0) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub speed: Option, + /// 音量(0-2,默认 1.0) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub volume: Option, + /// 音调(-1 到 1,默认 0) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub pitch: Option, + /// 情感风格(如 "happy"、"sad"、"neutral") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub emotion: Option, + /// 说话风格 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub style: Option, + /// 厂商特有参数透传 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +impl SpeechRequest { + /// 创建 Builder(单个音色) + pub fn builder( + model: impl Into, + input: impl Into, + voice: impl Into, + ) -> SpeechRequestBuilder { + Self::builder_with_voices(model, input, vec![voice.into()]) + } + + /// 创建 Builder(候选音色列表,用于自动降级) + pub fn builder_with_voices( + model: impl Into, + input: impl Into, + voices: Vec, + ) -> SpeechRequestBuilder { + SpeechRequestBuilder { + inner: SpeechRequest { + model: model.into(), + input: input.into(), + voice: VoiceSpec { voices }, + response_format: default_speech_format(), + speed: None, + volume: None, + pitch: None, + emotion: None, + style: None, + extra: HashMap::new(), + }, + } + } +} + +/// 音色规格(支持候选列表用于自动降级) +/// +/// 对应 Python v1 `speech` 的 `voice: str | list[str]` 参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct VoiceSpec { + /// 音色列表(至少 1 个;多个时启用 fallback 降级) + pub voices: Vec, +} + +impl VoiceSpec { + /// 单个音色 + pub fn single(v: impl Into) -> Self { + Self { + voices: vec![v.into()], + } + } + + /// 候选音色列表 + pub fn multiple(voices: Vec) -> Self { + Self { voices } + } + + /// 主音色(列表第一个) + pub fn primary(&self) -> Option<&str> { + self.voices.first().map(String::as_str) + } +} + +/// `SpeechRequest` 的 Builder +#[derive(Debug, Clone)] +pub struct SpeechRequestBuilder { + inner: SpeechRequest, +} + +impl SpeechRequestBuilder { + pub fn response_format(mut self, r: impl Into) -> Self { + self.inner.response_format = r.into(); + self + } + pub fn speed(mut self, s: f64) -> Self { + self.inner.speed = Some(s); + self + } + pub fn volume(mut self, v: f64) -> Self { + self.inner.volume = Some(v); + self + } + pub fn pitch(mut self, p: f64) -> Self { + self.inner.pitch = Some(p); + self + } + pub fn emotion(mut self, e: impl Into) -> Self { + self.inner.emotion = Some(e.into()); + self + } + pub fn style(mut self, s: impl Into) -> Self { + self.inner.style = Some(s.into()); + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> SpeechRequest { + self.inner + } +} + +/// 文字转语音结果 +/// +/// 对应 Python v1 `SpeechResult`。注意:`audio_data` 不参与 serde +/// (二进制数据通过 FFI 的 `aibridge_bytes_t` 单独传递),仅用于 core 内部。 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct SpeechResult { + /// 音频二进制数据(不序列化,FFI 单独传递) + #[serde(skip)] + pub audio_data: Option>, + /// 音频 URL(部分 Provider 返回) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub audio_url: Option, + /// 音频 Base64 编码 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub audio_base64: Option, + /// 音频 MIME 类型 + #[serde(default = "default_content_type")] + pub content_type: String, + /// 音频格式(mp3/wav/opus 等) + #[serde(default = "default_format")] + pub format: String, + /// 估计音频时长(秒) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub duration: Option, + /// 使用的模型 ID + #[serde(default, skip_serializing_if = "Option::is_none")] + pub model: Option, +} + +impl SpeechResult { + /// 获取音频二进制数据(优先 audio_data,其次解码 audio_base64) + pub fn get_audio_bytes(&self) -> Option> { + if let Some(data) = &self.audio_data { + return Some(data.clone()); + } + if let Some(b64) = &self.audio_base64 { + return crate::util::decode_base64(b64).ok(); + } + None + } +} + +fn default_transcribe_format() -> String { + "json".into() +} + +fn is_default_transcribe_format(s: &str) -> bool { + s == "json" +} + +fn default_speech_format() -> String { + "mp3".into() +} + +fn is_default_speech_format(s: &str) -> bool { + s == "mp3" +} + +fn default_task() -> String { + "transcribe".into() +} + +fn is_default_task(s: &str) -> bool { + s.is_empty() || s == "transcribe" +} + +fn default_content_type() -> String { + "audio/mpeg".into() +} + +fn default_format() -> String { + "mp3".into() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn transcribe_request_builder() { + let req = TranscribeRequest::builder("whisper-1", FileInput::path("/tmp/a.mp3")) + .language("zh") + .response_format("verbose_json") + .temperature(0.2) + .build(); + assert_eq!(req.model, "whisper-1"); + assert_eq!(req.language.as_deref(), Some("zh")); + assert_eq!(req.response_format, "verbose_json"); + } + + #[test] + fn transcribe_request_skip_defaults() { + let req = TranscribeRequest::builder("whisper-1", FileInput::path("/tmp/a.mp3")).build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(!json.contains("language")); + assert!(!json.contains("prompt")); + assert!(!json.contains("response_format")); // "json" 被跳过 + assert!(!json.contains("translate")); // false 被跳过 + } + + #[test] + fn transcribe_request_translate_flag() { + let req = TranscribeRequest::builder("whisper-1", FileInput::path("/tmp/a.mp3")) + .translate(true) + .build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(json.contains("\"translate\":true")); + } + + #[test] + fn transcription_result_deserialize() { + let json = serde_json::json!({ + "text": "hello world", + "language": "en", + "duration": 5.5, + "task": "transcribe" + }); + let r: TranscriptionResult = serde_json::from_value(json).unwrap(); + assert_eq!(r.text, "hello world"); + assert_eq!(r.language.as_deref(), Some("en")); + assert!((r.duration.unwrap() - 5.5).abs() < f64::EPSILON); + } + + #[test] + fn transcription_result_default_task() { + let r = TranscriptionResult { + text: "hi".into(), + ..Default::default() + }; + let json = serde_json::to_string(&r).unwrap(); + // 默认 "transcribe" 被跳过 + assert!(!json.contains("task")); + } + + #[test] + fn speech_request_builder_single_voice() { + let req = SpeechRequest::builder("tts-1", "hello", "alloy") + .speed(1.5) + .build(); + assert_eq!(req.model, "tts-1"); + assert_eq!(req.voice.primary(), Some("alloy")); + assert_eq!(req.speed, Some(1.5)); + } + + #[test] + fn speech_request_builder_voice_fallback() { + let req = SpeechRequest::builder_with_voices( + "edge-tts", + "hello", + vec!["zh-CN-XiaoxiaoNeural".into(), "zh-CN-YunxiNeural".into()], + ) + .build(); + assert_eq!(req.voice.voices.len(), 2); + assert_eq!(req.voice.primary(), Some("zh-CN-XiaoxiaoNeural")); + } + + #[test] + fn speech_request_skip_defaults() { + let req = SpeechRequest::builder("tts-1", "hi", "alloy").build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(!json.contains("response_format")); // "mp3" 被跳过 + assert!(!json.contains("speed")); + } + + #[test] + fn speech_result_get_audio_bytes_from_data() { + let r = SpeechResult { + audio_data: Some(vec![1, 2, 3]), + ..Default::default() + }; + assert_eq!(r.get_audio_bytes(), Some(vec![1, 2, 3])); + } + + #[test] + fn speech_result_get_audio_bytes_from_base64() { + let encoded = crate::util::encode_base64(b"hello"); + let r = SpeechResult { + audio_base64: Some(encoded), + ..Default::default() + }; + assert_eq!(r.get_audio_bytes(), Some(b"hello".to_vec())); + } + + #[test] + fn speech_result_get_audio_bytes_none() { + let r = SpeechResult::default(); + assert!(r.get_audio_bytes().is_none()); + } + + #[test] + fn speech_result_audio_data_not_serialized() { + let r = SpeechResult { + audio_data: Some(vec![1, 2, 3]), + ..Default::default() + }; + let json = serde_json::to_string(&r).unwrap(); + // audio_data 被 skip,不应出现 + assert!(!json.contains("audio_data")); + } + + #[test] + fn voice_spec_single_and_multiple() { + let s = VoiceSpec::single("alloy"); + assert_eq!(s.voices.len(), 1); + let m = VoiceSpec::multiple(vec!["a".into(), "b".into()]); + assert_eq!(m.voices.len(), 2); + } +} diff --git a/crates/aibridge-core/src/model/chat.rs b/crates/aibridge-core/src/model/chat.rs new file mode 100644 index 0000000..e1ff433 --- /dev/null +++ b/crates/aibridge-core/src/model/chat.rs @@ -0,0 +1,660 @@ +//! 文本对话数据模型 +//! +//! 定义文本对话相关的 serde struct:请求、消息、完成结果、流式块。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/chat.py`。 +//! +//! 设计要点(与设计文档第 6 节一致): +//! - `ChatMessage` 为 tagged enum(`role` 作为 tag),支持 system/user/assistant/tool +//! - `UserContent` 支持 String 或多模态 Vec +//! - `ChatRequest` 用 Builder 模式,替代 Python 的 `**kwargs` + `ChatOptions` + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +use crate::model::options::{ReasoningEffort, ResponseFormat, StopSeq, ToolChoice, ToolDefinition}; + +/// 对话请求 +/// +/// 对应设计文档第 6 节 `ChatRequest`。 +/// 用 `ChatRequest::builder(model, messages)` 链式构造。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatRequest { + /// 模型名称 + pub model: String, + /// 消息列表 + pub messages: Vec, + /// 温度系数(0-2,越高越随机) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub temperature: Option, + /// 核采样(0-1) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub top_p: Option, + /// Top-K 采样 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub top_k: Option, + /// 最大生成 token 数 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + /// 停止词 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub stop: Option, + /// 生成回复数量 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub n: Option, + /// 存在惩罚(-2 到 2) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub presence_penalty: Option, + /// 频率惩罚(-2 到 2) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub frequency_penalty: Option, + /// 随机种子 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub seed: Option, + /// 可用工具列表 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tools: Option>, + /// 工具选择策略 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_choice: Option, + /// 是否允许并行工具调用 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub parallel_tool_calls: Option, + /// 推理努力程度 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub reasoning_effort: Option, + /// 响应格式 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub response_format: Option, + /// 用户标识(用于风控/限流) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub user: Option, + /// 是否流式输出 + #[serde(default, skip_serializing_if = "std::ops::Not::not")] + pub stream: bool, + /// 厂商特有参数透传(不会被映射,直接加入请求体) + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +impl ChatRequest { + /// 创建 Builder + /// + /// # 示例 + /// ```ignore + /// let req = ChatRequest::builder("gpt-4o", vec![ + /// ChatMessage::user("Hello!"), + /// ]).temperature(0.7).max_tokens(1000).build(); + /// ``` + pub fn builder(model: impl Into, messages: Vec) -> ChatRequestBuilder { + ChatRequestBuilder { + inner: ChatRequest { + model: model.into(), + messages, + temperature: None, + top_p: None, + top_k: None, + max_tokens: None, + stop: None, + n: None, + presence_penalty: None, + frequency_penalty: None, + seed: None, + tools: None, + tool_choice: None, + parallel_tool_calls: None, + reasoning_effort: None, + response_format: None, + user: None, + stream: false, + extra: HashMap::new(), + }, + } + } +} + +/// `ChatRequest` 的 Builder +#[derive(Debug, Clone)] +pub struct ChatRequestBuilder { + inner: ChatRequest, +} + +impl ChatRequestBuilder { + pub fn temperature(mut self, t: f64) -> Self { + self.inner.temperature = Some(t); + self + } + pub fn top_p(mut self, t: f64) -> Self { + self.inner.top_p = Some(t); + self + } + pub fn top_k(mut self, t: u32) -> Self { + self.inner.top_k = Some(t); + self + } + pub fn max_tokens(mut self, t: u32) -> Self { + self.inner.max_tokens = Some(t); + self + } + pub fn stop(mut self, s: StopSeq) -> Self { + self.inner.stop = Some(s); + self + } + pub fn n(mut self, n: u32) -> Self { + self.inner.n = Some(n); + self + } + pub fn presence_penalty(mut self, p: f64) -> Self { + self.inner.presence_penalty = Some(p); + self + } + pub fn frequency_penalty(mut self, p: f64) -> Self { + self.inner.frequency_penalty = Some(p); + self + } + pub fn seed(mut self, s: u64) -> Self { + self.inner.seed = Some(s); + self + } + pub fn tools(mut self, t: Vec) -> Self { + self.inner.tools = Some(t); + self + } + pub fn tool_choice(mut self, c: ToolChoice) -> Self { + self.inner.tool_choice = Some(c); + self + } + pub fn parallel_tool_calls(mut self, p: bool) -> Self { + self.inner.parallel_tool_calls = Some(p); + self + } + pub fn reasoning_effort(mut self, e: ReasoningEffort) -> Self { + self.inner.reasoning_effort = Some(e); + self + } + pub fn response_format(mut self, f: ResponseFormat) -> Self { + self.inner.response_format = Some(f); + self + } + pub fn user(mut self, u: impl Into) -> Self { + self.inner.user = Some(u.into()); + self + } + pub fn stream(mut self, s: bool) -> Self { + self.inner.stream = s; + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> ChatRequest { + self.inner + } +} + +/// 对话消息(tagged enum,`role` 作为 tag) +/// +/// 对应设计文档第 6 节 `ChatMessage`。 +/// 对应 Python v1 `ChatMessage`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "role", rename_all = "lowercase")] +pub enum ChatMessage { + /// 系统消息 + System { + /// 消息内容 + content: String, + /// 发送者名称(可选) + #[serde(default, skip_serializing_if = "Option::is_none")] + name: Option, + }, + /// 用户消息 + User { + /// 消息内容(字符串或多模态) + content: UserContent, + /// 发送者名称(可选) + #[serde(default, skip_serializing_if = "Option::is_none")] + name: Option, + }, + /// 助手消息 + Assistant { + /// 消息内容 + #[serde(default, skip_serializing_if = "Option::is_none")] + content: Option, + /// 工具调用列表 + #[serde(default, skip_serializing_if = "Option::is_none")] + tool_calls: Option>, + }, + /// 工具结果消息 + Tool { + /// 工具调用 ID + tool_call_id: String, + /// 工具返回内容 + content: String, + }, +} + +impl ChatMessage { + /// 创建系统消息 + pub fn system(content: impl Into) -> Self { + Self::System { + content: content.into(), + name: None, + } + } + + /// 创建用户消息(纯文本) + pub fn user(content: impl Into) -> Self { + Self::User { + content: UserContent::Text(content.into()), + name: None, + } + } + + /// 创建用户消息(多模态) + pub fn user_multimodal(parts: Vec) -> Self { + Self::User { + content: UserContent::Parts(parts), + name: None, + } + } + + /// 创建助手消息 + pub fn assistant(content: impl Into) -> Self { + Self::Assistant { + content: Some(content.into()), + tool_calls: None, + } + } + + /// 创建工具结果消息 + pub fn tool(tool_call_id: impl Into, content: impl Into) -> Self { + Self::Tool { + tool_call_id: tool_call_id.into(), + content: content.into(), + } + } +} + +/// 用户消息内容(字符串或多模态部件列表) +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum UserContent { + /// 纯文本 + Text(String), + /// 多模态部件列表 + Parts(Vec), +} + +/// 多模态内容部件 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(tag = "type", rename_all = "snake_case")] +pub enum ContentPart { + /// 文本部件 + Text { + /// 文本内容 + text: String, + }, + /// 图片 URL 部件 + ImageUrl { + /// 图片 URL 信息 + image_url: ImageUrl, + }, +} + +/// 图片 URL 信息 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageUrl { + /// 图片 URL 或 data URI + pub url: String, + /// 图片细节级别("low" / "high" / "auto") + #[serde(default = "default_detail", skip_serializing_if = "is_default_detail")] + pub detail: String, +} + +impl ImageUrl { + /// 创建图片 URL + pub fn new(url: impl Into) -> Self { + Self { + url: url.into(), + detail: "auto".into(), + } + } + + /// 指定细节级别 + pub fn with_detail(mut self, detail: impl Into) -> Self { + self.detail = detail.into(); + self + } +} + +fn default_detail() -> String { + "auto".into() +} + +fn is_default_detail(d: &str) -> bool { + d == "auto" +} + +/// 对话完成结果 +/// +/// 对应 Python v1 `ChatCompletion`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatCompletion { + /// 响应 ID + pub id: String, + /// 对象类型 + #[serde(default = "default_object_completion")] + pub object: String, + /// 创建时间戳 + pub created: u64, + /// 使用的模型 + pub model: String, + /// 回复选项列表 + pub choices: Vec, + /// Token 使用统计 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, + /// 服务层级 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub service_tier: Option, + /// 系统指纹 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub system_fingerprint: Option, +} + +fn default_object_completion() -> String { + "chat.completion".into() +} + +/// 对话选项 +/// +/// 对应 Python v1 `ChatChoice`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatChoice { + /// 选项索引 + pub index: u32, + /// 生成的回复消息 + pub message: ChoiceMessage, + /// 结束原因(stop / length / content_filter / tool_calls) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub finish_reason: Option, +} + +/// 完成结果中的消息(比 ChatMessage 更宽松,便于解析 provider 响应) +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChoiceMessage { + /// 角色(通常为 "assistant") + pub role: String, + /// 消息内容 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + /// 工具调用列表 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +/// Token 使用统计 +/// +/// 对应 Python v1 `ChatUsage`。 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct ChatUsage { + /// 提示词 token 数 + pub prompt_tokens: u64, + /// 完成回复 token 数 + pub completion_tokens: u64, + /// 总 token 数 + pub total_tokens: u64, +} + +/// 流式对话块 +/// +/// 对应 Python v1 `ChatCompletionChunk`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatCompletionChunk { + /// 响应 ID + pub id: String, + /// 对象类型 + #[serde(default = "default_object_chunk")] + pub object: String, + /// 创建时间戳 + pub created: u64, + /// 使用的模型 + pub model: String, + /// 增量选项列表 + pub choices: Vec, + /// Token 使用统计(仅 stream_options.include_usage=true 时在末尾块出现) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +fn default_object_chunk() -> String { + "chat.completion.chunk".into() +} + +/// 流式增量 +/// +/// 对应 Python v1 `ChatCompletionDelta`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ChatCompletionDelta { + /// 增量索引 + pub index: u32, + /// 增量消息内容 + pub delta: DeltaMessage, + /// 结束原因 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub finish_reason: Option, +} + +/// 流式增量消息 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct DeltaMessage { + /// 角色(首个块通常为 "assistant") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub role: Option, + /// 增量内容 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub content: Option, + /// 工具调用增量 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub tool_calls: Option>, +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::options::{FunctionDefinition, ToolCallFunction}; + + #[test] + fn chat_message_system_serde() { + let msg = ChatMessage::system("You are helpful."); + let json = serde_json::to_string(&msg).unwrap(); + assert!(json.contains("\"role\":\"system\"")); + assert!(json.contains("\"content\":\"You are helpful.\"")); + let back: ChatMessage = serde_json::from_str(&json).unwrap(); + match back { + ChatMessage::System { content, .. } => assert_eq!(content, "You are helpful."), + _ => panic!("应为 System"), + } + } + + #[test] + fn chat_message_user_text_serde() { + let msg = ChatMessage::user("Hi"); + let json = serde_json::to_string(&msg).unwrap(); + assert!(json.contains("\"role\":\"user\"")); + // UserContent::Text 序列化为字符串 + assert!(json.contains("\"content\":\"Hi\"")); + } + + #[test] + fn chat_message_user_multimodal_serde() { + let msg = ChatMessage::user_multimodal(vec![ + ContentPart::Text { + text: "What's this?".into(), + }, + ContentPart::ImageUrl { + image_url: ImageUrl::new("https://example.com/img.png"), + }, + ]); + let json = serde_json::to_string(&msg).unwrap(); + assert!(json.contains("\"role\":\"user\"")); + assert!(json.contains("\"type\":\"text\"")); + assert!(json.contains("\"type\":\"image_url\"")); + assert!(json.contains("https://example.com/img.png")); + } + + #[test] + fn chat_message_assistant_with_tool_calls() { + let msg = ChatMessage::Assistant { + content: None, + tool_calls: Some(vec![crate::model::options::ToolCall { + id: "call_1".into(), + tool_type: "function".into(), + function: ToolCallFunction { + name: "get_weather".into(), + arguments: "{}".into(), + }, + }]), + }; + let json = serde_json::to_string(&msg).unwrap(); + assert!(json.contains("\"role\":\"assistant\"")); + assert!(json.contains("\"tool_calls\"")); + } + + #[test] + fn chat_message_tool_serde() { + let msg = ChatMessage::tool("call_1", "sunny"); + let json = serde_json::to_string(&msg).unwrap(); + assert!(json.contains("\"role\":\"tool\"")); + assert!(json.contains("\"tool_call_id\":\"call_1\"")); + assert!(json.contains("\"content\":\"sunny\"")); + } + + #[test] + fn chat_request_builder() { + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(1000) + .top_p(0.9) + .stream(true) + .extra("custom", "value") + .build(); + assert_eq!(req.model, "gpt-4o"); + assert_eq!(req.temperature, Some(0.7)); + assert_eq!(req.max_tokens, Some(1000)); + assert!(req.stream); + assert_eq!( + req.extra.get("custom").and_then(|v| v.as_str()), + Some("value") + ); + } + + #[test] + fn chat_request_skip_none_fields() { + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(!json.contains("temperature")); + assert!(!json.contains("max_tokens")); + assert!(!json.contains("stream")); // false 且 skip_serializing_if + assert!(!json.contains("extra")); // 空且 skip_serializing_if + } + + #[test] + fn chat_request_with_tools() { + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("weather?")]) + .tools(vec![ToolDefinition::function(FunctionDefinition { + name: "get_weather".into(), + description: None, + parameters: None, + })]) + .tool_choice(ToolChoice::Auto) + .build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(json.contains("\"tools\"")); + assert!(json.contains("\"tool_choice\":\"auto\"")); + } + + #[test] + fn chat_completion_deserialize_openai_format() { + let json = serde_json::json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": "Hello!" + }, + "finish_reason": "stop" + }], + "usage": { + "prompt_tokens": 5, + "completion_tokens": 2, + "total_tokens": 7 + } + }); + let comp: ChatCompletion = serde_json::from_value(json).unwrap(); + assert_eq!(comp.id, "chatcmpl-1"); + assert_eq!(comp.model, "gpt-4o"); + assert_eq!(comp.choices.len(), 1); + assert_eq!(comp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(comp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(comp.usage.as_ref().unwrap().total_tokens, 7); + } + + #[test] + fn chat_completion_chunk_deserialize() { + let json = serde_json::json!({ + "id": "chatcmpl-1", + "object": "chat.completion.chunk", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "delta": {"content": "Hello"}, + "finish_reason": null + }] + }); + let chunk: ChatCompletionChunk = serde_json::from_value(json).unwrap(); + assert_eq!(chunk.choices[0].delta.content.as_deref(), Some("Hello")); + assert!(chunk.choices[0].finish_reason.is_none()); + } + + #[test] + fn image_url_default_detail_auto() { + let iu = ImageUrl::new("https://example.com/x.png"); + assert_eq!(iu.detail, "auto"); + let json = serde_json::to_string(&iu).unwrap(); + // detail == "auto" 时被 skip + assert!(!json.contains("detail")); + } + + #[test] + fn image_url_custom_detail_kept() { + let iu = ImageUrl::new("https://example.com/x.png").with_detail("high"); + let json = serde_json::to_string(&iu).unwrap(); + assert!(json.contains("\"detail\":\"high\"")); + } + + #[test] + fn user_content_text_roundtrip() { + let c = UserContent::Text("hello".into()); + let json = serde_json::to_string(&c).unwrap(); + assert_eq!(json, "\"hello\""); + let back: UserContent = serde_json::from_str(&json).unwrap(); + match back { + UserContent::Text(s) => assert_eq!(s, "hello"), + _ => panic!("应为 Text"), + } + } + + #[test] + fn delta_message_default_empty() { + let d = DeltaMessage::default(); + assert!(d.role.is_none()); + assert!(d.content.is_none()); + } +} diff --git a/crates/aibridge-core/src/model/common.rs b/crates/aibridge-core/src/model/common.rs new file mode 100644 index 0000000..0d1daa1 --- /dev/null +++ b/crates/aibridge-core/src/model/common.rs @@ -0,0 +1,385 @@ +//! 通用数据模型 +//! +//! 定义模型类型常量、模型信息、Provider 信息等通用数据结构。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/common.py`。 + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +/// 模型类型 +/// +/// 对应 Python v1 `ModelType` 常量类。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ModelType { + /// 文本对话 + Chat, + /// 图像生成 + Image, + /// 视频生成 + Video, + /// 音频(ASR/TTS) + Audio, +} + +impl ModelType { + /// 转为字符串 + pub fn as_str(&self) -> &'static str { + match self { + Self::Chat => "chat", + Self::Image => "image", + Self::Video => "video", + Self::Audio => "audio", + } + } +} + +/// 从字符串解析模型类型 +impl From<&str> for ModelType { + fn from(s: &str) -> Self { + match s.to_lowercase().as_str() { + "image" => Self::Image, + "video" => Self::Video, + "audio" => Self::Audio, + _ => Self::Chat, + } + } +} + +/// 视频生成模式 +/// +/// 对应 Python v1 `VideoMode` 常量类。 +#[derive(Debug, Clone, Copy, Default, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum VideoMode { + /// 文生视频 + #[default] + Text2Video, + /// 图生视频 + Image2Video, + /// 关键帧模式 + Keyframes, + /// 多图模式 + Multiimage, +} + +/// 任务状态 +/// +/// 对应 Python v1 `TaskStatus` 常量类。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum TaskStatus { + /// 排队中 + Pending, + /// 处理中 + Processing, + /// 成功 + Success, + /// 失败 + Failed, +} + +impl TaskStatus { + /// 是否终态 + pub fn is_terminal(&self) -> bool { + matches!(self, Self::Success | Self::Failed) + } +} + +/// 模型信息 +/// +/// 描述单个 AI 模型的基本信息和能力。 +/// 对应 Python v1 `ModelInfo`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ModelInfo { + /// 模型标识符 + pub id: String, + /// 模型显示名称 + pub name: String, + /// 模型类型 + #[serde(rename = "type")] + pub model_type: ModelType, + /// 提供商名称 + pub provider: String, + /// 支持的能力列表(如 "text2image"、"image2image") + #[serde(default)] + pub capabilities: Vec, + /// 最大 token 数(仅 chat 模型) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub max_tokens: Option, + /// 是否支持流式输出 + #[serde(default)] + pub supports_streaming: bool, + /// 模型描述 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub description: Option, + /// 模型创建时间戳 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub created: Option, +} + +/// Provider 信息 +/// +/// 描述单个 AI 模型提供商的元信息。 +/// 对应 Python v1 `ProviderInfo`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ProviderInfo { + /// Provider 类型标识 + #[serde(rename = "type")] + pub provider_type: String, + /// Provider 显示名称 + pub name: String, + /// Provider 描述 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub description: Option, + /// Provider 官网 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub website: Option, + /// 支持的能力列表 + #[serde(default)] + pub supported_capabilities: Vec, + /// 支持的模型类型 + #[serde(default)] + pub supported_model_types: Vec, +} + +/// 语音信息 +/// +/// 用于 TTS Provider 的音色描述(edge-tts / elevenlabs 等返回)。 +/// 字段较为宽松,以适配不同 Provider 的音色元数据格式。 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct VoiceInfo { + /// 音色短名(edge-tts 的 ShortName / elevenlabs 的 voice_id) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub short_name: Option, + /// 音色显示名 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub name: Option, + /// 语言区域(如 "zh-CN" / "en-US") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub locale: Option, + /// 性别("Female" / "Male") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub gender: Option, + /// 音色 ID(部分 Provider 用) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub voice_id: Option, + /// 额外元数据 + #[serde(default, flatten)] + pub extra: HashMap, +} + +impl VoiceInfo { + /// 创建一个新的 Builder + pub fn builder() -> VoiceInfoBuilder { + VoiceInfoBuilder::default() + } +} + +/// `VoiceInfo` 的 Builder +#[derive(Debug, Default, Clone)] +pub struct VoiceInfoBuilder { + inner: VoiceInfo, +} + +impl VoiceInfoBuilder { + pub fn short_name(mut self, v: impl Into) -> Self { + self.inner.short_name = Some(v.into()); + self + } + pub fn name(mut self, v: impl Into) -> Self { + self.inner.name = Some(v.into()); + self + } + pub fn locale(mut self, v: impl Into) -> Self { + self.inner.locale = Some(v.into()); + self + } + pub fn gender(mut self, v: impl Into) -> Self { + self.inner.gender = Some(v.into()); + self + } + pub fn voice_id(mut self, v: impl Into) -> Self { + self.inner.voice_id = Some(v.into()); + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> VoiceInfo { + self.inner + } +} + +/// 模型类型推断关键字(用于从 /models 端点拉取的模型 ID 推断类型) +/// +/// 对应 Python v1 `BaseAdapter._infer_type` 的关键字表。 +pub fn infer_model_type(model_id: &str) -> ModelType { + let lower = model_id.to_lowercase(); + let image_keywords = [ + "image", + "flux", + "sd3", + "sdxl", + "dall", + "seedream", + "wanx", + "ideogram", + "midjourney", + "stable-diffusion", + "imagen", + ]; + let video_keywords = [ + "video", + "veo", + "seedance", + "cogvideox", + "wan", + "kling", + "runway", + "pika", + "luma", + "sora", + "vidu", + ]; + let audio_keywords = [ + "whisper", + "tts", + "speech", + "transcribe", + "nova", + "sonic", + "edge-tts", + "eleven", + "cosyvoice", + "sensevoice", + ]; + if image_keywords.iter().any(|kw| lower.contains(kw)) { + ModelType::Image + } else if video_keywords.iter().any(|kw| lower.contains(kw)) { + ModelType::Video + } else if audio_keywords.iter().any(|kw| lower.contains(kw)) { + ModelType::Audio + } else { + ModelType::Chat + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn model_type_serde_lowercase() { + let json = serde_json::to_string(&ModelType::Image).unwrap(); + assert_eq!(json, "\"image\""); + let t: ModelType = serde_json::from_str("\"video\"").unwrap(); + assert_eq!(t, ModelType::Video); + } + + #[test] + fn model_type_from_str() { + assert_eq!(ModelType::from("chat"), ModelType::Chat); + assert_eq!(ModelType::from("IMAGE"), ModelType::Image); + assert_eq!(ModelType::from("unknown"), ModelType::Chat); + } + + #[test] + fn model_type_as_str() { + assert_eq!(ModelType::Audio.as_str(), "audio"); + } + + #[test] + fn video_mode_default() { + assert_eq!(VideoMode::default(), VideoMode::Text2Video); + } + + #[test] + fn task_status_is_terminal() { + assert!(TaskStatus::Success.is_terminal()); + assert!(TaskStatus::Failed.is_terminal()); + assert!(!TaskStatus::Pending.is_terminal()); + assert!(!TaskStatus::Processing.is_terminal()); + } + + #[test] + fn infer_model_type_image() { + assert_eq!(infer_model_type("dall-e-3"), ModelType::Image); + assert_eq!(infer_model_type("seedream-4.0"), ModelType::Image); + assert_eq!(infer_model_type("FLUX-schnell"), ModelType::Image); + } + + #[test] + fn infer_model_type_video() { + assert_eq!(infer_model_type("seedance-2.0"), ModelType::Video); + assert_eq!(infer_model_type("kling-v1"), ModelType::Video); + assert_eq!(infer_model_type("veo-3"), ModelType::Video); + } + + #[test] + fn infer_model_type_audio() { + assert_eq!(infer_model_type("whisper-1"), ModelType::Audio); + assert_eq!(infer_model_type("tts-1-hd"), ModelType::Audio); + assert_eq!(infer_model_type("edge-tts"), ModelType::Audio); + } + + #[test] + fn infer_model_type_chat_default() { + assert_eq!(infer_model_type("gpt-4o"), ModelType::Chat); + assert_eq!(infer_model_type("claude-3-opus"), ModelType::Chat); + } + + #[test] + fn model_info_serialize() { + let m = ModelInfo { + id: "gpt-4o".into(), + name: "GPT-4o".into(), + model_type: ModelType::Chat, + provider: "openai".into(), + capabilities: vec!["chat".into()], + max_tokens: Some(128000), + supports_streaming: true, + description: None, + created: None, + }; + let json = serde_json::to_string(&m).unwrap(); + assert!(json.contains("\"type\":\"chat\"")); + assert!(json.contains("\"max_tokens\":128000")); + // skip_serializing_if 生效:description/created 不出现 + assert!(!json.contains("description")); + assert!(!json.contains("created")); + } + + #[test] + fn voice_info_builder() { + let v = VoiceInfo::builder() + .short_name("zh-CN-XiaoxiaoNeural") + .name("Xiaoxiao") + .locale("zh-CN") + .gender("Female") + .extra("wheels", 4) + .build(); + assert_eq!(v.short_name.as_deref(), Some("zh-CN-XiaoxiaoNeural")); + assert_eq!(v.gender.as_deref(), Some("Female")); + assert_eq!(v.extra.get("wheels").and_then(|x| x.as_i64()), Some(4)); + } + + #[test] + fn voice_info_flatten_extra() { + let v = VoiceInfo { + short_name: Some("v1".into()), + ..Default::default() + }; + let json = serde_json::json!({"short_name": "v1", "custom": "x"}); + let parsed: VoiceInfo = serde_json::from_value(json).unwrap(); + assert_eq!(parsed.short_name.as_deref(), Some("v1")); + assert_eq!( + parsed.extra.get("custom").and_then(|x| x.as_str()), + Some("x") + ); + // 序列化回带 extra + let _ = v; + } +} diff --git a/crates/aibridge-core/src/model/image.rs b/crates/aibridge-core/src/model/image.rs new file mode 100644 index 0000000..39a5b5a --- /dev/null +++ b/crates/aibridge-core/src/model/image.rs @@ -0,0 +1,382 @@ +//! 图像生成数据模型 +//! +//! 定义图像生成相关的 serde struct。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/image.py`。 +//! +//! 设计要点(与设计文档第 6 节一致): +//! - 去掉 Python 的 `ImageOptions` 中间层,改为 `ImageRequest::builder()` +//! - `FileInput` 抽象图像输入(Path/Url/Bytes/Base64),适配多 provider + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +/// 图像生成请求 +/// +/// 对应设计文档第 6 节,合并 Python v1 `ImageGenerationOptions` + 请求参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageRequest { + /// 模型名称 + pub model: String, + /// 提示词 + pub prompt: String, + /// 图像尺寸(如 "1024x1024"),或使用 width/height + #[serde(default, skip_serializing_if = "Option::is_none")] + pub size: Option, + /// 宽度(像素) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub width: Option, + /// 高度(像素) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub height: Option, + /// 画面比例(如 "16:9") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub aspect_ratio: Option, + /// 生成数量(1-10) + #[serde(default = "default_n", skip_serializing_if = "is_default_n")] + pub n: u32, + /// 生成质量("standard" / "hd" / "ultra") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub quality: Option, + /// 图像风格 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub style: Option, + /// 负面提示词 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub negative_prompt: Option, + /// 随机种子 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub seed: Option, + /// 推理步数 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub steps: Option, + /// CFG Scale(提示词相关性) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cfg_scale: Option, + /// 响应格式("url" / "b64_json") + #[serde( + default = "default_response_format", + skip_serializing_if = "is_default_response_format" + )] + pub response_format: String, + /// 输出图片格式("png" / "jpeg" / "webp") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub output_format: Option, + /// 参考图(图生图/IP-Adapter),FileInput 列表 + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub reference_images: Vec, + /// 遮罩图片(局部重绘) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub mask: Option, + /// 编辑模式(inpaint / outpaint / variation) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub edit_mode: Option, + /// 厂商特有参数透传 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +impl ImageRequest { + /// 创建 Builder + pub fn builder(model: impl Into, prompt: impl Into) -> ImageRequestBuilder { + ImageRequestBuilder { + inner: ImageRequest { + model: model.into(), + prompt: prompt.into(), + size: None, + width: None, + height: None, + aspect_ratio: None, + n: default_n(), + quality: None, + style: None, + negative_prompt: None, + seed: None, + steps: None, + cfg_scale: None, + response_format: default_response_format(), + output_format: None, + reference_images: Vec::new(), + mask: None, + edit_mode: None, + extra: HashMap::new(), + }, + } + } +} + +/// `ImageRequest` 的 Builder +#[derive(Debug, Clone)] +pub struct ImageRequestBuilder { + inner: ImageRequest, +} + +impl ImageRequestBuilder { + pub fn size(mut self, s: impl Into) -> Self { + self.inner.size = Some(s.into()); + self + } + pub fn width(mut self, w: u32) -> Self { + self.inner.width = Some(w); + self + } + pub fn height(mut self, h: u32) -> Self { + self.inner.height = Some(h); + self + } + pub fn aspect_ratio(mut self, a: impl Into) -> Self { + self.inner.aspect_ratio = Some(a.into()); + self + } + pub fn n(mut self, n: u32) -> Self { + self.inner.n = n; + self + } + pub fn quality(mut self, q: impl Into) -> Self { + self.inner.quality = Some(q.into()); + self + } + pub fn style(mut self, s: impl Into) -> Self { + self.inner.style = Some(s.into()); + self + } + pub fn negative_prompt(mut self, n: impl Into) -> Self { + self.inner.negative_prompt = Some(n.into()); + self + } + pub fn seed(mut self, s: u64) -> Self { + self.inner.seed = Some(s); + self + } + pub fn steps(mut self, s: u32) -> Self { + self.inner.steps = Some(s); + self + } + pub fn cfg_scale(mut self, c: f64) -> Self { + self.inner.cfg_scale = Some(c); + self + } + pub fn response_format(mut self, r: impl Into) -> Self { + self.inner.response_format = r.into(); + self + } + pub fn output_format(mut self, o: impl Into) -> Self { + self.inner.output_format = Some(o.into()); + self + } + pub fn reference_images(mut self, imgs: Vec) -> Self { + self.inner.reference_images = imgs; + self + } + pub fn mask(mut self, m: FileInput) -> Self { + self.inner.mask = Some(m); + self + } + pub fn edit_mode(mut self, e: impl Into) -> Self { + self.inner.edit_mode = Some(e.into()); + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> ImageRequest { + self.inner + } +} + +/// 文件输入 +/// +/// 对应设计文档第 6 节 `FileInput`。 +/// 抽象各种图像输入形式(路径、URL、字节、Base64)。 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum FileInput { + /// 本地文件路径 + Path(String), + /// 远程 URL + Url(String), + /// 原始字节 + Bytes(Vec), + /// Base64 编码字符串 + Base64(String), +} + +impl FileInput { + /// 从 URL 创建 + pub fn url(u: impl Into) -> Self { + Self::Url(u.into()) + } + + /// 从路径创建 + pub fn path(p: impl Into) -> Self { + Self::Path(p.into()) + } + + /// 从 Base64 创建 + pub fn base64(b: impl Into) -> Self { + Self::Base64(b.into()) + } + + /// 从字节创建 + pub fn bytes(b: Vec) -> Self { + Self::Bytes(b) + } +} + +/// 图像数据 +/// +/// 对应 Python v1 `ImageData`。 +#[derive(Debug, Clone, Serialize, Deserialize, Default)] +pub struct ImageData { + /// 图像 URL + #[serde(default, skip_serializing_if = "Option::is_none")] + pub url: Option, + /// Base64 编码的图像 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub b64_json: Option, + /// 修改后的提示词(如模型优化过) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub revised_prompt: Option, +} + +/// 图像生成结果 +/// +/// 对应 Python v1 `ImageGenerationResult`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ImageResult { + /// 响应 ID + pub id: String, + /// 对象类型 + #[serde(default = "default_object_image")] + pub object: String, + /// 创建时间戳 + pub created: u64, + /// 使用的模型 + pub model: String, + /// 生成的图像列表 + pub data: Vec, +} + +fn default_object_image() -> String { + "image.generation".into() +} + +fn default_n() -> u32 { + 1 +} + +fn is_default_n(n: &u32) -> bool { + *n == 1 +} + +fn default_response_format() -> String { + "url".into() +} + +fn is_default_response_format(s: &str) -> bool { + s == "url" +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn image_request_builder_defaults() { + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + assert_eq!(req.model, "dall-e-3"); + assert_eq!(req.prompt, "a cat"); + assert_eq!(req.n, 1); + assert_eq!(req.response_format, "url"); + } + + #[test] + fn image_request_builder_chained() { + let req = ImageRequest::builder("dall-e-3", "a cat") + .size("1024x1024") + .n(2) + .quality("hd") + .style("vivid") + .seed(42) + .build(); + assert_eq!(req.size.as_deref(), Some("1024x1024")); + assert_eq!(req.n, 2); + assert_eq!(req.quality.as_deref(), Some("hd")); + assert_eq!(req.seed, Some(42)); + } + + #[test] + fn image_request_skip_defaults() { + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(!json.contains("\"n\"")); // n=1 被跳过 + assert!(!json.contains("\"response_format\"")); // "url" 被跳过 + assert!(!json.contains("negative_prompt")); + assert!(!json.contains("extra")); + } + + #[test] + fn image_request_with_reference_images() { + let req = ImageRequest::builder("flux", "edit this") + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .mask(FileInput::base64("aGVsbG8=")) + .edit_mode("inpaint") + .build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(json.contains("\"reference_images\"")); + assert!(json.contains("\"mask\"")); + assert!(json.contains("\"edit_mode\":\"inpaint\"")); + } + + #[test] + fn file_input_url_serde() { + let f = FileInput::url("https://example.com/x.png"); + let json = serde_json::to_string(&f).unwrap(); + assert_eq!(json, "\"https://example.com/x.png\""); + } + + #[test] + fn file_input_base64_serde() { + let f = FileInput::base64("aGVsbG8="); + let json = serde_json::to_string(&f).unwrap(); + assert_eq!(json, "\"aGVsbG8=\""); + } + + #[test] + fn image_result_deserialize() { + let json = serde_json::json!({ + "id": "img-1", + "object": "image.generation", + "created": 1700000000, + "model": "dall-e-3", + "data": [{ + "url": "https://example.com/result.png", + "revised_prompt": "a cute cat" + }] + }); + let r: ImageResult = serde_json::from_value(json).unwrap(); + assert_eq!(r.id, "img-1"); + assert_eq!(r.model, "dall-e-3"); + assert_eq!(r.data.len(), 1); + assert_eq!( + r.data[0].url.as_deref(), + Some("https://example.com/result.png") + ); + } + + #[test] + fn image_data_default_empty() { + let d = ImageData::default(); + assert!(d.url.is_none()); + assert!(d.b64_json.is_none()); + } + + #[test] + fn file_input_constructors() { + let _ = FileInput::path("/tmp/x.png"); + let _ = FileInput::bytes(vec![1, 2, 3]); + let _ = FileInput::url("https://x"); + let _ = FileInput::base64("aGk="); + } +} diff --git a/crates/aibridge-core/src/model/mod.rs b/crates/aibridge-core/src/model/mod.rs new file mode 100644 index 0000000..0b0e757 --- /dev/null +++ b/crates/aibridge-core/src/model/mod.rs @@ -0,0 +1,40 @@ +//! 数据模型层 +//! +//! 定义所有 AI 能力的 serde struct,替代 Python v1 的 Pydantic 模型。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/` 目录。 +//! +//! 子模块: +//! - [`common`]:通用类型(ModelType / ModelInfo / VoiceInfo / TaskStatus 等) +//! - [`options`]:工具定义、嵌入类型、参数映射等公共类型 +//! - [`chat`]:文本对话(ChatRequest / ChatMessage / ChatCompletion / ChatCompletionChunk) +//! - [`image`]:图像生成(ImageRequest / ImageResult / FileInput) +//! - [`video`]:视频生成(VideoRequest / VideoTask / VideoStatus) +//! - [`audio`]:语音(TranscribeRequest / TranscriptionResult / SpeechRequest / SpeechResult) + +pub mod audio; +pub mod chat; +pub mod common; +pub mod image; +pub mod options; +pub mod video; + +// 重新导出常用类型,方便上层使用 +pub use audio::{ + SpeechRequest, SpeechRequestBuilder, SpeechResult, TranscribeRequest, TranscribeRequestBuilder, + TranscriptionResult, TranscriptionSegment, TranscriptionWord, VoiceSpec, +}; +pub use chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChatRequestBuilder, ChatUsage, ChoiceMessage, ContentPart, DeltaMessage, ImageUrl, UserContent, +}; +pub use common::{ + infer_model_type, ModelInfo, ModelType, ProviderInfo, TaskStatus, VideoMode, VoiceInfo, + VoiceInfoBuilder, +}; +pub use image::{FileInput, ImageData, ImageRequest, ImageRequestBuilder, ImageResult}; +pub use options::{ + EmbedInput, EmbedRequest, EmbeddingItem, EmbeddingResult, EmbeddingUsage, EmbeddingVector, + FunctionDefinition, ParameterMapping, ReasoningEffort, ResponseFormat, StopSeq, ToolCall, + ToolCallFunction, ToolChoice, ToolDefinition, +}; +pub use video::{VideoRequest, VideoRequestBuilder, VideoStatus, VideoTask}; diff --git a/crates/aibridge-core/src/model/options.rs b/crates/aibridge-core/src/model/options.rs new file mode 100644 index 0000000..9220462 --- /dev/null +++ b/crates/aibridge-core/src/model/options.rs @@ -0,0 +1,416 @@ +//! 统一请求选项与工具定义 +//! +//! 定义所有 AI 能力的标准化请求参数及工具/函数调用相关类型。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/options.py`。 +//! +//! 设计原则(与设计文档第 6 节一致): +//! - Rust 不保留 Python 的 `ChatOptions/ImageOptions` 中间层,改为 `Request::builder()` 链式调用 +//! - 厂商特有参数通过 `extra: HashMap` 透传 +//! - 工具调用、响应格式等通用类型在此统一定义 + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +/// 推理努力程度(统一思考模式) +/// +/// 对应 Python v1 `ReasoningEffort`。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ReasoningEffort { + None, + Low, + Medium, + High, + Auto, +} + +/// 响应格式 +/// +/// 对应 Python v1 `ResponseFormat`。 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "snake_case")] +pub enum ResponseFormat { + /// 纯文本 + Text, + /// JSON 对象 + JsonObject, + /// JSON Schema(结构化输出) + JsonSchema, +} + +/// 工具选择策略 +/// +/// 对应 Python v1 `ToolChoice`。 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(rename_all = "lowercase")] +pub enum ToolChoice { + /// 不调用工具 + None, + /// 自动决定 + Auto, + /// 强制调用 + Required, +} + +/// 停止词 +/// +/// 单个或多个停止词。 +#[derive(Debug, Clone, PartialEq, Eq, Serialize, Deserialize)] +#[serde(untagged)] +pub enum StopSeq { + /// 单个停止词 + Single(String), + /// 多个停止词 + Multiple(Vec), +} + +/// 函数定义 +/// +/// 对应 Python v1 `FunctionDefinition`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct FunctionDefinition { + /// 函数名称 + pub name: String, + /// 函数描述 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub description: Option, + /// 函数参数(JSON Schema) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub parameters: Option, +} + +/// 工具定义 +/// +/// 对应 Python v1 `ToolDefinition`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolDefinition { + /// 工具类型 + #[serde(default = "default_tool_type")] + #[serde(rename = "type")] + pub tool_type: String, + /// 函数定义(type=function 时必填) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub function: Option, +} + +impl ToolDefinition { + /// 创建函数类型工具 + pub fn function(func: FunctionDefinition) -> Self { + Self { + tool_type: "function".into(), + function: Some(func), + } + } + + /// 创建 web_search 工具 + pub fn web_search() -> Self { + Self { + tool_type: "web_search".into(), + function: None, + } + } +} + +fn default_tool_type() -> String { + "function".into() +} + +/// 工具调用(模型生成的调用请求) +/// +/// 对应 Python v1 `ToolCall`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolCall { + /// 工具调用 ID + pub id: String, + /// 工具类型(通常为 "function") + #[serde(default = "default_tool_type")] + #[serde(rename = "type")] + pub tool_type: String, + /// 函数调用信息:包含 `name` 和 `arguments`(arguments 为 JSON 字符串) + pub function: ToolCallFunction, +} + +/// 工具调用中的函数部分 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct ToolCallFunction { + /// 函数名 + pub name: String, + /// 函数参数(JSON 字符串,由模型生成) + pub arguments: String, +} + +/// 通用参数映射规则 +/// +/// 对应 Python v1 `ParameterMapping`。 +/// 定义通用参数名 → 厂商特定参数名的映射关系。 +/// +/// 阶段 1 起由适配器使用;此处仅定义结构,预置常量在适配器层。 +#[derive(Debug, Clone, Default)] +pub struct ParameterMapping { + /// 键名重命名表(值为 None 表示移除该参数) + pub rename_map: HashMap>, +} + +impl ParameterMapping { + /// 创建空映射 + pub fn new() -> Self { + Self::default() + } + + /// 应用映射规则 + pub fn apply( + &self, + params: &HashMap, + ) -> HashMap { + let mut result = HashMap::new(); + for (key, value) in params { + match self.rename_map.get(key) { + // 显式移除 + Some(None) => continue, + // 重命名 + Some(Some(new_key)) => { + result.insert(new_key.clone(), value.clone()); + } + // 保持原名 + None => { + result.insert(key.clone(), value.clone()); + } + } + } + result + } +} + +/// 文本嵌入请求 +/// +/// 对应 Python v1 `EmbedOptions` + 嵌入请求参数。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct EmbedRequest { + /// 嵌入模型名称 + pub model: String, + /// 输入文本(单个或多个) + pub input: EmbedInput, + /// 输出向量维度(部分模型支持) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub dimensions: Option, + /// 编码格式("float" / "base64") + #[serde(default, skip_serializing_if = "Option::is_none")] + pub encoding_format: Option, + /// 用户标识 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub user: Option, + /// 厂商特有参数透传 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +/// 嵌入输入(单个字符串或字符串列表) +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum EmbedInput { + /// 单个文本 + Single(String), + /// 多个文本 + Multiple(Vec), +} + +/// 嵌入结果 +/// +/// 对应 Python v1 `EmbeddingResult`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct EmbeddingResult { + /// 对象类型(固定 "list") + #[serde(default = "default_object_list")] + pub object: String, + /// 嵌入向量列表 + pub data: Vec, + /// 使用的模型 + pub model: String, + /// 使用统计 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub usage: Option, +} + +/// 单个嵌入项 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct EmbeddingItem { + /// 对象类型(固定 "embedding") + #[serde(default = "default_object_embedding")] + pub object: String, + /// 索引 + pub index: u32, + /// 嵌入向量(浮点列表,或 base64 字符串) + pub embedding: EmbeddingVector, +} + +/// 嵌入向量表示 +#[derive(Debug, Clone, Serialize, Deserialize)] +#[serde(untagged)] +pub enum EmbeddingVector { + /// 浮点列表 + Float(Vec), + /// base64 编码 + Base64(String), +} + +/// 嵌入使用统计 +#[derive(Debug, Clone, Default, Serialize, Deserialize)] +pub struct EmbeddingUsage { + /// 提示词 token 数 + pub prompt_tokens: u64, + /// 总 token 数 + pub total_tokens: u64, +} + +fn default_object_list() -> String { + "list".into() +} + +fn default_object_embedding() -> String { + "embedding".into() +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn reasoning_effort_serde() { + let json = serde_json::to_string(&ReasoningEffort::High).unwrap(); + assert_eq!(json, "\"high\""); + } + + #[test] + fn response_format_serde() { + let json = serde_json::to_string(&ResponseFormat::JsonObject).unwrap(); + assert_eq!(json, "\"json_object\""); + } + + #[test] + fn tool_choice_serde() { + let json = serde_json::to_string(&ToolChoice::Auto).unwrap(); + assert_eq!(json, "\"auto\""); + } + + #[test] + fn stop_seq_untagged() { + let single = serde_json::to_string(&StopSeq::Single("stop".into())).unwrap(); + assert_eq!(single, "\"stop\""); + let multi = + serde_json::to_string(&StopSeq::Multiple(vec!["a".into(), "b".into()])).unwrap(); + assert_eq!(multi, "[\"a\",\"b\"]"); + } + + #[test] + fn tool_definition_function() { + let td = ToolDefinition::function(FunctionDefinition { + name: "get_weather".into(), + description: Some("Get weather".into()), + parameters: None, + }); + let json = serde_json::to_string(&td).unwrap(); + assert!(json.contains("\"type\":\"function\"")); + assert!(json.contains("\"name\":\"get_weather\"")); + } + + #[test] + fn tool_definition_web_search() { + let td = ToolDefinition::web_search(); + let json = serde_json::to_string(&td).unwrap(); + assert!(json.contains("\"type\":\"web_search\"")); + } + + #[test] + fn tool_call_serde() { + let tc = ToolCall { + id: "call_1".into(), + tool_type: "function".into(), + function: ToolCallFunction { + name: "get_weather".into(), + arguments: "{\"city\":\"Beijing\"}".into(), + }, + }; + let json = serde_json::to_string(&tc).unwrap(); + assert!(json.contains("\"id\":\"call_1\"")); + assert!(json.contains("\"arguments\":\"{\\\"city\\\":\\\"Beijing\\\"}\"")); + } + + #[test] + fn parameter_mapping_rename() { + let mut rename = HashMap::new(); + rename.insert("max_tokens".into(), Some("maxOutputTokens".into())); + rename.insert("web_search".into(), None); // 移除 + let pm = ParameterMapping { rename_map: rename }; + + let mut params = HashMap::new(); + params.insert("max_tokens".into(), serde_json::json!(1000)); + params.insert("web_search".into(), serde_json::json!(true)); + params.insert("temperature".into(), serde_json::json!(0.7)); + + let result = pm.apply(¶ms); + assert_eq!( + result.get("maxOutputTokens").and_then(|v| v.as_i64()), + Some(1000) + ); + assert!(!result.contains_key("max_tokens")); + assert!(!result.contains_key("web_search")); + assert_eq!( + result.get("temperature").and_then(|v| v.as_f64()), + Some(0.7) + ); + } + + #[test] + fn embed_input_serde() { + let single = EmbedInput::Single("hello".into()); + let json = serde_json::to_string(&single).unwrap(); + assert_eq!(json, "\"hello\""); + + let multi = EmbedInput::Multiple(vec!["a".into(), "b".into()]); + let json = serde_json::to_string(&multi).unwrap(); + assert_eq!(json, "[\"a\",\"b\"]"); + } + + #[test] + fn embedding_result_serde() { + let r = EmbeddingResult { + object: "list".into(), + data: vec![EmbeddingItem { + object: "embedding".into(), + index: 0, + embedding: EmbeddingVector::Float(vec![0.1, 0.2, 0.3]), + }], + model: "text-embedding-3-small".into(), + usage: Some(EmbeddingUsage { + prompt_tokens: 5, + total_tokens: 5, + }), + }; + let json = serde_json::to_string(&r).unwrap(); + assert!(json.contains("\"object\":\"list\"")); + assert!(json.contains("\"model\":\"text-embedding-3-small\"")); + } + + #[test] + fn embedding_vector_base64() { + let v = EmbeddingVector::Base64("aGVsbG8=".into()); + let json = serde_json::to_string(&v).unwrap(); + assert_eq!(json, "\"aGVsbG8=\""); + } + + #[test] + fn embed_request_skip_empty_extra() { + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let json = serde_json::to_string(&req).unwrap(); + assert!(!json.contains("extra")); + assert!(!json.contains("dimensions")); + } +} diff --git a/crates/aibridge-core/src/model/video.rs b/crates/aibridge-core/src/model/video.rs new file mode 100644 index 0000000..f7f3924 --- /dev/null +++ b/crates/aibridge-core/src/model/video.rs @@ -0,0 +1,385 @@ +//! 视频生成数据模型 +//! +//! 定义视频生成相关的 serde struct。 +//! 对应 Python v1 (agn-sdk) 的 `agn/models/video.py`。 + +use std::collections::HashMap; + +use serde::{Deserialize, Serialize}; + +use crate::model::common::{TaskStatus, VideoMode}; +use crate::model::image::FileInput; + +/// 视频生成请求 +/// +/// 对应设计文档第 6 节,合并 Python v1 `VideoGenerationOptions` + `VideoTaskCreate`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct VideoRequest { + /// 模型名称 + pub model: String, + /// 提示词 + pub prompt: String, + /// 视频宽度(必须是 8 的倍数) + #[serde(default = "default_width", skip_serializing_if = "is_default_width")] + pub width: u32, + /// 视频高度(必须是 8 的倍数) + #[serde(default = "default_height", skip_serializing_if = "is_default_height")] + pub height: u32, + /// 帧数(部分模型需要) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub num_frames: Option, + /// 帧率 + #[serde( + default = "default_frame_rate", + skip_serializing_if = "is_default_frame_rate" + )] + pub frame_rate: u32, + /// 生成模式 + #[serde(default)] + pub mode: VideoMode, + /// 视频时长(秒),部分模型用此字段替代 num_frames + #[serde(default, skip_serializing_if = "Option::is_none")] + pub duration: Option, + /// 宽高比,如 "16:9" / "9:16" / "1:1" + #[serde(default, skip_serializing_if = "Option::is_none")] + pub aspect_ratio: Option, + /// 分辨率档位,如 "720p" / "1080p" + #[serde(default, skip_serializing_if = "Option::is_none")] + pub resolution: Option, + /// 参考图像列表(图生视频) + #[serde(default, skip_serializing_if = "Vec::is_empty")] + pub reference_images: Vec, + /// 首帧图片 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub first_frame: Option, + /// 尾帧图片 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub last_frame: Option, + /// 镜头运动 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub camera_motion: Option, + /// 运动强度(0-10) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub motion_strength: Option, + /// 负面提示词 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub negative_prompt: Option, + /// 随机种子 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub seed: Option, + /// 推理步数 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub steps: Option, + /// CFG Scale + #[serde(default, skip_serializing_if = "Option::is_none")] + pub cfg_scale: Option, + /// 是否生成音频 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub with_audio: Option, + /// 是否添加水印 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub watermark: Option, + /// 厂商特有参数透传 + #[serde(default, skip_serializing_if = "HashMap::is_empty")] + pub extra: HashMap, +} + +impl VideoRequest { + /// 创建 Builder + pub fn builder(model: impl Into, prompt: impl Into) -> VideoRequestBuilder { + VideoRequestBuilder { + inner: VideoRequest { + model: model.into(), + prompt: prompt.into(), + width: default_width(), + height: default_height(), + num_frames: None, + frame_rate: default_frame_rate(), + mode: VideoMode::default(), + duration: None, + aspect_ratio: None, + resolution: None, + reference_images: Vec::new(), + first_frame: None, + last_frame: None, + camera_motion: None, + motion_strength: None, + negative_prompt: None, + seed: None, + steps: None, + cfg_scale: None, + with_audio: None, + watermark: None, + extra: HashMap::new(), + }, + } + } +} + +/// `VideoRequest` 的 Builder +#[derive(Debug, Clone)] +pub struct VideoRequestBuilder { + inner: VideoRequest, +} + +impl VideoRequestBuilder { + pub fn width(mut self, w: u32) -> Self { + self.inner.width = w; + self + } + pub fn height(mut self, h: u32) -> Self { + self.inner.height = h; + self + } + pub fn num_frames(mut self, n: u32) -> Self { + self.inner.num_frames = Some(n); + self + } + pub fn frame_rate(mut self, f: u32) -> Self { + self.inner.frame_rate = f; + self + } + pub fn mode(mut self, m: VideoMode) -> Self { + self.inner.mode = m; + self + } + pub fn duration(mut self, d: u32) -> Self { + self.inner.duration = Some(d); + self + } + pub fn aspect_ratio(mut self, a: impl Into) -> Self { + self.inner.aspect_ratio = Some(a.into()); + self + } + pub fn resolution(mut self, r: impl Into) -> Self { + self.inner.resolution = Some(r.into()); + self + } + pub fn reference_images(mut self, imgs: Vec) -> Self { + self.inner.reference_images = imgs; + self + } + pub fn first_frame(mut self, f: FileInput) -> Self { + self.inner.first_frame = Some(f); + self + } + pub fn last_frame(mut self, f: FileInput) -> Self { + self.inner.last_frame = Some(f); + self + } + pub fn camera_motion(mut self, c: impl Into) -> Self { + self.inner.camera_motion = Some(c.into()); + self + } + pub fn motion_strength(mut self, m: f64) -> Self { + self.inner.motion_strength = Some(m); + self + } + pub fn negative_prompt(mut self, n: impl Into) -> Self { + self.inner.negative_prompt = Some(n.into()); + self + } + pub fn seed(mut self, s: u64) -> Self { + self.inner.seed = Some(s); + self + } + pub fn steps(mut self, s: u32) -> Self { + self.inner.steps = Some(s); + self + } + pub fn cfg_scale(mut self, c: f64) -> Self { + self.inner.cfg_scale = Some(c); + self + } + pub fn with_audio(mut self, w: bool) -> Self { + self.inner.with_audio = Some(w); + self + } + pub fn watermark(mut self, w: bool) -> Self { + self.inner.watermark = Some(w); + self + } + pub fn extra(mut self, k: impl Into, v: impl Into) -> Self { + self.inner.extra.insert(k.into(), v.into()); + self + } + pub fn build(self) -> VideoRequest { + self.inner + } +} + +/// 视频任务信息(创建任务后的返回) +/// +/// 对应 Python v1 `VideoTask`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct VideoTask { + /// 任务 ID(用于轮询状态) + pub task_id: String, + /// 使用的模型 + pub model: String, + /// 任务状态 + #[serde(default = "default_status")] + pub status: TaskStatus, + /// 创建时间戳 + pub created_at: u64, +} + +/// 视频任务状态(轮询返回) +/// +/// 对应 Python v1 `VideoStatus`。 +#[derive(Debug, Clone, Serialize, Deserialize)] +pub struct VideoStatus { + /// 任务 ID + pub task_id: String, + /// 任务状态 + pub status: TaskStatus, + /// 视频 URL(成功时) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub video_url: Option, + /// 进度 0-100 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub progress: Option, + /// 错误信息(失败时) + #[serde(default, skip_serializing_if = "Option::is_none")] + pub error: Option, + /// 创建时间戳 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub created_at: Option, + /// 更新时间戳 + #[serde(default, skip_serializing_if = "Option::is_none")] + pub updated_at: Option, +} + +fn default_width() -> u32 { + 1280 +} + +fn is_default_width(w: &u32) -> bool { + *w == 1280 +} + +fn default_height() -> u32 { + 720 +} + +fn is_default_height(h: &u32) -> bool { + *h == 720 +} + +fn default_frame_rate() -> u32 { + 24 +} + +fn is_default_frame_rate(f: &u32) -> bool { + *f == 24 +} + +fn default_status() -> TaskStatus { + TaskStatus::Pending +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn video_request_builder_defaults() { + let req = VideoRequest::builder("seedance-2.0", "a cat walking").build(); + assert_eq!(req.model, "seedance-2.0"); + assert_eq!(req.width, 1280); + assert_eq!(req.height, 720); + assert_eq!(req.frame_rate, 24); + assert_eq!(req.mode, VideoMode::Text2Video); + } + + #[test] + fn video_request_builder_chained() { + let req = VideoRequest::builder("seedance-2.0", "a cat") + .width(1920) + .height(1080) + .duration(5) + .aspect_ratio("16:9") + .seed(123) + .with_audio(true) + .build(); + assert_eq!(req.width, 1920); + assert_eq!(req.height, 1080); + assert_eq!(req.duration, Some(5)); + assert_eq!(req.with_audio, Some(true)); + } + + #[test] + fn video_request_skip_defaults() { + let req = VideoRequest::builder("seedance-2.0", "a cat").build(); + let json = serde_json::to_string(&req).unwrap(); + // 默认值被跳过 + assert!(!json.contains("\"width\"")); + assert!(!json.contains("\"height\"")); + assert!(!json.contains("\"frame_rate\"")); + assert!(!json.contains("negative_prompt")); + assert!(!json.contains("extra")); + } + + #[test] + fn video_request_image2video_mode() { + let req = VideoRequest::builder("kling-v1", "animate this") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let json = serde_json::to_string(&req).unwrap(); + assert!(json.contains("\"mode\":\"image2video\"")); + assert!(json.contains("\"reference_images\"")); + } + + #[test] + fn video_task_deserialize() { + let json = serde_json::json!({ + "task_id": "t-1", + "model": "seedance-2.0", + "status": "pending", + "created_at": 1700000000 + }); + let t: VideoTask = serde_json::from_value(json).unwrap(); + assert_eq!(t.task_id, "t-1"); + assert_eq!(t.model, "seedance-2.0"); + assert_eq!(t.status, TaskStatus::Pending); + } + + #[test] + fn video_status_success() { + let json = serde_json::json!({ + "task_id": "t-1", + "status": "success", + "video_url": "https://example.com/v.mp4", + "progress": 100 + }); + let s: VideoStatus = serde_json::from_value(json).unwrap(); + assert_eq!(s.status, TaskStatus::Success); + assert_eq!(s.video_url.as_deref(), Some("https://example.com/v.mp4")); + assert_eq!(s.progress, Some(100)); + } + + #[test] + fn video_status_failed_with_error() { + let json = serde_json::json!({ + "task_id": "t-1", + "status": "failed", + "error": "content policy violation" + }); + let s: VideoStatus = serde_json::from_value(json).unwrap(); + assert_eq!(s.status, TaskStatus::Failed); + assert_eq!(s.error.as_deref(), Some("content policy violation")); + } + + #[test] + fn video_status_processing_with_progress() { + let json = serde_json::json!({ + "task_id": "t-1", + "status": "processing", + "progress": 45 + }); + let s: VideoStatus = serde_json::from_value(json).unwrap(); + assert_eq!(s.status, TaskStatus::Processing); + assert_eq!(s.progress, Some(45)); + } +} diff --git a/crates/aibridge-core/src/retry.rs b/crates/aibridge-core/src/retry.rs new file mode 100644 index 0000000..a7422d5 --- /dev/null +++ b/crates/aibridge-core/src/retry.rs @@ -0,0 +1,306 @@ +//! 重试机制 +//! +//! 提供异步重试工具,支持指数退避策略。 +//! 对应 Python v1 (agn-sdk) 的 `agn/core/retry.py`(替代 tenacity)。 +//! +//! 设计要点: +//! - 仅对可重试错误重试(`AibridgeError::is_retryable`) +//! - 指数退避:`delay * multiplier^attempt`,封顶 `max_delay` +//! - 限流错误优先采用服务端 `retry_after` + +use std::future::Future; +use std::time::Duration; + +use tokio::time::sleep; + +use crate::error::{AibridgeError, Result}; + +/// 重试策略配置 +/// +/// 对应 Python v1 `retry_on_error` 的参数。 +#[derive(Debug, Clone, Copy)] +pub struct RetryPolicy { + /// 最大尝试次数(含首次) + pub max_attempts: u32, + + /// 初始延迟(秒) + pub delay: f64, + + /// 退避倍数 + pub multiplier: f64, + + /// 最大延迟(秒) + pub max_delay: f64, +} + +impl Default for RetryPolicy { + fn default() -> Self { + Self { + max_attempts: 3, + delay: 2.0, + multiplier: 2.0, + max_delay: 60.0, + } + } +} + +impl RetryPolicy { + /// 创建一个新的 Builder + pub fn builder() -> RetryPolicyBuilder { + RetryPolicyBuilder::default() + } + + /// 计算第 `attempt` 次失败后的等待时长(attempt 从 1 开始) + /// + /// 公式:`min(delay * multiplier^(attempt-1), max_delay)`。 + /// 若错误携带 `retry_after`,优先使用 `retry_after`。 + pub fn backoff(&self, attempt: u32) -> Duration { + if attempt == 0 { + return Duration::ZERO; + } + let exp = self.multiplier.powi((attempt - 1) as i32); + let secs = (self.delay * exp).min(self.max_delay).max(0.0); + Duration::from_secs_f64(secs) + } + + /// 计算针对特定错误的等待时长(限流错误优先用 retry_after) + pub fn backoff_for(&self, attempt: u32, err: &AibridgeError) -> Duration { + if let Some(retry_after) = err.retry_after() { + // 服务端建议值封顶到 max_delay,避免过长等待 + let capped = if retry_after > Duration::from_secs_f64(self.max_delay) { + Duration::from_secs_f64(self.max_delay) + } else { + retry_after + }; + // 至少为指数退避值(取较大者,尊重服务端的下限建议) + capped.max(self.backoff(attempt)) + } else { + self.backoff(attempt) + } + } +} + +/// `RetryPolicy` 的 Builder +#[derive(Debug, Default, Clone, Copy)] +pub struct RetryPolicyBuilder { + inner: RetryPolicy, +} + +impl RetryPolicyBuilder { + pub fn max_attempts(mut self, n: u32) -> Self { + self.inner.max_attempts = n; + self + } + + pub fn delay(mut self, delay: f64) -> Self { + self.inner.delay = delay; + self + } + + pub fn multiplier(mut self, multiplier: f64) -> Self { + self.inner.multiplier = multiplier; + self + } + + pub fn max_delay(mut self, max_delay: f64) -> Self { + self.inner.max_delay = max_delay; + self + } + + pub fn build(self) -> RetryPolicy { + self.inner + } +} + +/// 对异步操作执行重试 +/// +/// 对应 Python v1 `retry_async`。仅当返回的错误可重试时重试。 +/// +/// # 参数 +/// - `policy`: 重试策略 +/// - `operation`: 异步闭包,返回 `Result` +/// +/// # 示例 +/// ```ignore +/// let result: Result = retry_with(&policy, || async { +/// // 可能失败的异步操作 +/// Ok(42) +/// }).await; +/// ``` +pub async fn retry_with(policy: &RetryPolicy, operation: F) -> Result +where + F: Fn() -> Fut, + Fut: Future>, +{ + let mut last_err: Option = None; + for attempt in 1..=policy.max_attempts { + match operation().await { + Ok(v) => return Ok(v), + Err(e) => { + let retryable = e.is_retryable(); + last_err = Some(e); + if !retryable || attempt == policy.max_attempts { + break; + } + let wait = match last_err { + Some(ref err) => policy.backoff_for(attempt, err), + None => policy.backoff(attempt), + }; + if !wait.is_zero() { + sleep(wait).await; + } + } + } + } + Err(last_err.unwrap_or_else(|| AibridgeError::Api { + status: 0, + message: "retry_with: 未产生错误但流程异常".into(), + })) +} + +#[cfg(test)] +mod tests { + use super::*; + use std::sync::atomic::{AtomicU32, Ordering}; + use std::sync::Arc; + + #[test] + fn backoff_grows_exponentially() { + let p = RetryPolicy { + max_attempts: 5, + delay: 1.0, + multiplier: 2.0, + max_delay: 100.0, + }; + assert_eq!(p.backoff(1), Duration::from_secs(1)); + assert_eq!(p.backoff(2), Duration::from_secs(2)); + assert_eq!(p.backoff(3), Duration::from_secs(4)); + assert_eq!(p.backoff(4), Duration::from_secs(8)); + } + + #[test] + fn backoff_capped_at_max_delay() { + let p = RetryPolicy { + max_attempts: 10, + delay: 1.0, + multiplier: 2.0, + max_delay: 5.0, + }; + // 第 4 次:1*2^3 = 8,封顶到 5 + assert_eq!(p.backoff(4), Duration::from_secs(5)); + } + + #[test] + fn backoff_zero_for_attempt_zero() { + let p = RetryPolicy::default(); + assert_eq!(p.backoff(0), Duration::ZERO); + } + + #[test] + fn backoff_for_rate_limit_uses_retry_after() { + let p = RetryPolicy { + max_attempts: 5, + delay: 1.0, + multiplier: 2.0, + max_delay: 60.0, + }; + let err = AibridgeError::rate_limit_with_retry("slow", 3.0); + assert_eq!(p.backoff_for(1, &err), Duration::from_secs(3)); + } + + #[test] + fn backoff_for_rate_limit_capped_by_max_delay() { + let p = RetryPolicy { + max_attempts: 5, + delay: 1.0, + multiplier: 2.0, + max_delay: 5.0, + }; + let err = AibridgeError::rate_limit_with_retry("slow", 100.0); + assert_eq!(p.backoff_for(1, &err), Duration::from_secs(5)); + } + + #[test] + fn backoff_for_non_rate_limit_uses_exponential() { + let p = RetryPolicy { + max_attempts: 5, + delay: 1.0, + multiplier: 2.0, + max_delay: 60.0, + }; + let err = AibridgeError::Timeout; + assert_eq!(p.backoff_for(2, &err), Duration::from_secs(2)); + } + + #[tokio::test] + async fn retry_succeeds_on_first_attempt() { + let policy = RetryPolicy::default(); + let result: Result = retry_with(&policy, || async { Ok(42) }).await; + assert_eq!(result.unwrap(), 42); + } + + #[tokio::test] + async fn retry_succeeds_after_transient_failure() { + let policy = RetryPolicy::builder() + .max_attempts(3) + .delay(0.001) + .multiplier(1.0) + .max_delay(0.01) + .build(); + let counter = Arc::new(AtomicU32::new(0)); + let c = counter.clone(); + let result: Result = retry_with(&policy, || { + let c = c.clone(); + async move { + let n = c.fetch_add(1, Ordering::SeqCst); + if n < 1 { + Err(AibridgeError::Timeout) + } else { + Ok(7) + } + } + }) + .await; + assert_eq!(result.unwrap(), 7); + assert_eq!(counter.load(Ordering::SeqCst), 2); + } + + #[tokio::test] + async fn retry_gives_up_after_max_attempts() { + let policy = RetryPolicy::builder() + .max_attempts(2) + .delay(0.001) + .max_delay(0.01) + .build(); + let result: Result = + retry_with(&policy, || async { Err(AibridgeError::Timeout) }).await; + assert!(matches!(result, Err(AibridgeError::Timeout))); + } + + #[tokio::test] + async fn retry_does_not_retry_non_retryable_error() { + let policy = RetryPolicy::builder().max_attempts(5).delay(0.001).build(); + let counter = Arc::new(AtomicU32::new(0)); + let c = counter.clone(); + let result: Result = retry_with(&policy, || { + let c = c.clone(); + async move { + c.fetch_add(1, Ordering::SeqCst); + Err(AibridgeError::authentication("bad key")) + } + }) + .await; + // 非可重试错误应立即返回,不重试 + assert!(matches!(result, Err(AibridgeError::Authentication { .. }))); + assert_eq!(counter.load(Ordering::SeqCst), 1); + } + + #[test] + fn builder_defaults() { + let p = RetryPolicy::builder().build(); + assert_eq!(p.max_attempts, 3); + assert!((p.delay - 2.0).abs() < f64::EPSILON); + assert!((p.multiplier - 2.0).abs() < f64::EPSILON); + assert!((p.max_delay - 60.0).abs() < f64::EPSILON); + } +} diff --git a/crates/aibridge-core/src/router.rs b/crates/aibridge-core/src/router.rs new file mode 100644 index 0000000..2f4af0a --- /dev/null +++ b/crates/aibridge-core/src/router.rs @@ -0,0 +1,899 @@ +//! 路由器 +//! +//! 支持多 Provider 路由、负载均衡、Fallback。 +//! 对应 Python v1 (agn-sdk) 的 `agn/router.py`。 +//! +//! 设计要点: +//! - 多 Provider 配置,按策略选择(first / round_robin / random / weighted) +//! - 模型名 → Provider 的映射表(迁移自 Python v1 `MODEL_PROVIDER_MAP`) +//! - Fallback:主 Provider 失败时切换备用 +//! - 健康状态跟踪 + 延迟统计 +//! - 适配器用 `Arc` 存储,从锁中克隆后立即释放锁再调用(避免跨 await 持锁) +//! +//! 阶段 0.4 注意:因具体适配器未实现,`start()` 创建适配器会失败, +//! 实际路由能力在阶段 1 适配器就绪后可用。本模块的结构与逻辑已完整, +//! 单测用 mock adapter 验证路由选择与 fallback 行为。 + +use std::collections::HashMap; +use std::sync::{Arc, RwLock}; +use std::time::Instant; + +use rand::Rng; + +use crate::adapter::{create_adapter, Adapter, Capabilities, ChatStream}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; +use crate::model::{ + ChatCompletion, ChatRequest, EmbedRequest, EmbeddingResult, ImageRequest, ImageResult, + SpeechRequest, SpeechResult, TranscribeRequest, TranscriptionResult, VideoRequest, VideoStatus, + VideoTask, +}; + +/// 路由策略 +/// +/// 对应 Python v1 `RoutingStrategy`。 +#[derive(Debug, Clone, Copy, PartialEq, Eq, Default)] +pub enum RoutingStrategy { + /// 按顺序选第一个可用的 + #[default] + First, + /// 轮询 + RoundRobin, + /// 随机 + Random, + /// 按权重随机 + Weighted, + /// 按延迟(无数据时回退到 First) + Latency, +} + +/// 单个 Provider 的路由配置 +#[derive(Debug, Clone)] +pub struct ProviderEntry { + /// Provider 类型 + pub provider_type: String, + /// 连接选项 + pub options: ClientOptions, + /// 权重(weighted 策略用,默认 1) + pub weight: u32, +} + +impl ProviderEntry { + /// 创建一个 Provider 路由配置 + pub fn new(provider_type: impl Into, options: ClientOptions) -> Self { + Self { + provider_type: provider_type.into(), + options, + weight: 1, + } + } + + /// 设置权重 + pub fn with_weight(mut self, w: u32) -> Self { + self.weight = w.max(1); + self + } +} + +/// 多 Provider 路由器 +/// +/// 对应 Python v1 `Router`。 +pub struct Router { + entries: Vec, + default_provider: Option, + strategy: RoutingStrategy, + enable_fallback: bool, + max_retries: u32, + + // 运行时状态(用 RwLock 支持 &self 方法) + inner: RwLock, +} + +#[derive(Default)] +struct RouterInner { + /// 已启动的适配器(provider_type -> Arc) + adapters: HashMap>, + /// provider 顺序 + provider_order: Vec, + /// 健康状态 + health: HashMap, + /// 延迟统计(秒) + latency: HashMap, + /// round_robin 计数器 + rr_index: usize, + /// 自定义模型映射 + model_map: HashMap, + /// 权重表(从 entries 拷贝,便于策略选择时读取) + weights: HashMap, +} + +impl Router { + /// 创建路由器 + pub fn new(entries: Vec) -> Self { + Self::with_strategy(entries, RoutingStrategy::default()) + } + + /// 创建路由器(指定策略) + pub fn with_strategy(entries: Vec, strategy: RoutingStrategy) -> Self { + let weights: HashMap = entries + .iter() + .map(|e| (e.provider_type.clone(), e.weight)) + .collect(); + Self { + entries, + default_provider: None, + strategy, + enable_fallback: true, + max_retries: 2, + inner: RwLock::new(RouterInner { + weights, + ..Default::default() + }), + } + } + + /// 设置默认 Provider + pub fn with_default_provider(mut self, provider: impl Into) -> Self { + self.default_provider = Some(provider.into()); + self + } + + /// 设置是否启用 fallback + pub fn with_fallback(mut self, enable: bool) -> Self { + self.enable_fallback = enable; + self + } + + /// 设置 fallback 最大重试次数 + pub fn with_max_retries(mut self, n: u32) -> Self { + self.max_retries = n; + self + } + + /// 启动路由器(初始化所有适配器) + /// + /// 单个 Provider 启动失败不会中断整体,仅标记为不健康。 + pub async fn start(&self) -> Result<()> { + for entry in &self.entries { + let config = + ProviderConfig::from_options(entry.provider_type.clone(), entry.options.clone()); + match create_adapter(config) { + Ok(mut adapter) => { + if let Err(e) = adapter.start().await { + tracing::warn!( + provider = %entry.provider_type, + error = %e, + "启动 Provider 失败" + ); + let mut inner = self.inner.write().unwrap(); + inner.health.insert(entry.provider_type.clone(), false); + } else { + let mut inner = self.inner.write().unwrap(); + inner + .adapters + .insert(entry.provider_type.clone(), Arc::from(adapter)); + inner.provider_order.push(entry.provider_type.clone()); + inner.health.insert(entry.provider_type.clone(), true); + inner.latency.insert(entry.provider_type.clone(), 0.0); + } + } + Err(e) => { + tracing::warn!( + provider = %entry.provider_type, + error = %e, + "创建 Provider 适配器失败" + ); + let mut inner = self.inner.write().unwrap(); + inner.health.insert(entry.provider_type.clone(), false); + } + } + } + Ok(()) + } + + /// 关闭路由器(释放所有资源) + /// + /// 注意:`Arc` 可能被多个持有者共享,close 仅尝试通知; + /// 真正的资源释放在最后一个 Arc drop 时发生。 + pub async fn close(&self) -> Result<()> { + let adapters = { + let mut inner = self.inner.write().unwrap(); + let ads: Vec> = inner.adapters.values().cloned().collect(); + inner.adapters.clear(); + inner.provider_order.clear(); + inner.health.clear(); + inner.latency.clear(); + ads + }; + // Arc 无法直接 close(需要 &mut),这里只能依赖 drop。 + // 为保持 API 一致,显式 drop。 + drop(adapters); + Ok(()) + } + + /// 注册自定义模型映射 + pub fn register_model_mapping(&self, model: impl Into, provider: impl Into) { + let mut inner = self.inner.write().unwrap(); + inner.model_map.insert(model.into(), provider.into()); + } + + /// 获取所有 Provider 的健康状态 + pub fn get_health_status(&self) -> HashMap { + self.inner.read().unwrap().health.clone() + } + + /// 获取所有 Provider 的延迟统计 + pub fn get_latency_stats(&self) -> HashMap { + self.inner.read().unwrap().latency.clone() + } + + /// 选择 Provider(按模型名 + 能力) + fn select_provider(&self, model: &str, capability: Capabilities) -> Result { + let inner = self.inner.read().unwrap(); + + // 1. 自定义映射优先 + if let Some(p) = inner.model_map.get(model) { + if inner.adapters.contains_key(p) { + return Ok(p.clone()); + } + } + // 2. 内置映射 + if let Some(p) = builtin_model_provider(model) { + if inner.adapters.contains_key(p) { + return Ok(p.to_string()); + } + } + // 3. 按能力筛选候选 + let candidates = capable_providers(&inner, capability); + if candidates.is_empty() { + // 4. 默认 Provider + if let Some(ref dp) = self.default_provider { + if inner.adapters.contains_key(dp) { + return Ok(dp.clone()); + } + } + return Err(AibridgeError::model_not_found(format!( + "无法为模型 '{model}' 找到合适的 Provider(能力:{})", + capability.as_str() + ))); + } + Ok(pick_by_strategy(&candidates, self.strategy, &inner)) + } + + /// 获取 fallback 候选(排除已失败的) + fn fallback_providers(&self, failed: &str, capability: Capabilities) -> Vec { + let inner = self.inner.read().unwrap(); + let candidates = capable_providers(&inner, capability); + candidates.into_iter().filter(|p| p != failed).collect() + } + + /// 从锁中取出指定 provider 的 Arc clone(不在锁内调用) + fn get_adapter(&self, provider: &str) -> Option> { + self.inner.read().unwrap().adapters.get(provider).cloned() + } + + /// 带Fallback 执行 + async fn execute_with_fallback( + &self, + model: &str, + capability: Capabilities, + op: F, + ) -> Result + where + F: Fn(Arc) -> Fut, + Fut: std::future::Future>, + { + let primary = self.select_provider(model, capability)?; + let mut to_try = vec![primary.clone()]; + if self.enable_fallback { + let fb = self.fallback_providers(&primary, capability); + to_try.extend(fb.into_iter().take(self.max_retries as usize)); + } + + let mut last_err: Option = None; + for (i, provider_type) in to_try.iter().enumerate() { + let Some(adapter) = self.get_adapter(provider_type) else { + continue; + }; + + let start = Instant::now(); + match op(adapter).await { + Ok(v) => { + let elapsed = start.elapsed().as_secs_f64(); + let mut inner = self.inner.write().unwrap(); + let prev = inner.latency.get(provider_type).copied().unwrap_or(0.0); + let new = if prev > 0.0 { + prev * 0.7 + elapsed * 0.3 + } else { + elapsed + }; + inner.latency.insert(provider_type.clone(), new); + inner.health.insert(provider_type.clone(), true); + if i > 0 { + tracing::info!(from = %primary, to = %provider_type, "Fallback 成功"); + } + return Ok(v); + } + Err(e) => { + tracing::warn!( + provider = %provider_type, + attempt = i + 1, + error = %e, + "Provider 调用失败" + ); + last_err = Some(e); + let mut inner = self.inner.write().unwrap(); + inner.health.insert(provider_type.clone(), false); + } + } + } + Err(last_err.unwrap_or_else(|| AibridgeError::Api { + status: 0, + message: "所有 Provider 均失败".into(), + })) + } + + /// 文本对话 + pub async fn chat(&self, req: ChatRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::Chat, |adapter| { + let req = req.clone(); + async move { adapter.chat(req).await } + }) + .await + } + + /// 流式文本对话(不支持 fallback) + pub async fn chat_stream(&self, req: ChatRequest) -> Result { + let model = req.model.clone(); + let provider = self.select_provider(&model, Capabilities::ChatStream)?; + let Some(adapter) = self.get_adapter(&provider) else { + return Err(AibridgeError::model_not_found(format!( + "Provider '{provider}' 未启动" + ))); + }; + adapter.chat_stream(req).await + } + + /// 图像生成 + pub async fn image_generate(&self, req: ImageRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::ImageGenerate, |adapter| { + let req = req.clone(); + async move { adapter.image_generate(req).await } + }) + .await + } + + /// 创建视频生成任务 + pub async fn video_create(&self, req: VideoRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::VideoGenerate, |adapter| { + let req = req.clone(); + async move { adapter.video_create(req).await } + }) + .await + } + + /// 查询视频任务状态 + pub async fn video_poll(&self, task_id: &str, model: &str) -> Result { + // 视频轮询:按模型映射找 provider,否则遍历支持 video 的 + let provider = { + let inner = self.inner.read().unwrap(); + if let Some(p) = inner.model_map.get(model) { + Some(p.clone()) + } else { + builtin_model_provider(model).map(|p| p.to_string()) + } + }; + + if let Some(p) = provider { + if let Some(adapter) = self.get_adapter(&p) { + return adapter.video_poll(task_id, model).await; + } + } + + // 遍历支持 video 的 provider + let providers: Vec = { + let inner = self.inner.read().unwrap(); + capable_providers(&inner, Capabilities::VideoGenerate) + }; + for p in providers { + if let Some(adapter) = self.get_adapter(&p) { + match adapter.video_poll(task_id, model).await { + Ok(v) => return Ok(v), + Err(_) => continue, + } + } + } + Err(AibridgeError::model_not_found(format!( + "无法确定视频轮询的 Provider(task_id={task_id}, model={model})" + ))) + } + + /// 文本嵌入 + pub async fn embed(&self, req: EmbedRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::Embedding, |adapter| { + let req = req.clone(); + async move { adapter.embed(req).await } + }) + .await + } + + /// 语音转文字 + pub async fn transcribe(&self, req: TranscribeRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::AudioTranscribe, |adapter| { + let req = req.clone(); + async move { adapter.transcribe(req).await } + }) + .await + } + + /// 文字转语音 + pub async fn speech(&self, req: SpeechRequest) -> Result { + let model = req.model.clone(); + self.execute_with_fallback(&model, Capabilities::AudioSpeech, |adapter| { + let req = req.clone(); + async move { adapter.speech(req).await } + }) + .await + } + + /// 获取可用模型列表(聚合所有 Provider) + pub async fn list_models(&self, filter: Option) -> Result> { + let providers: Vec = { + let inner = self.inner.read().unwrap(); + inner + .provider_order + .iter() + .filter(|p| *inner.health.get(*p).unwrap_or(&true)) + .cloned() + .collect() + }; + + let mut all = Vec::new(); + let mut seen = std::collections::HashSet::new(); + for p in providers { + if let Some(adapter) = self.get_adapter(&p) { + match adapter.list_models(filter).await { + Ok(models) => { + for m in models { + if seen.insert(m.id.clone()) { + all.push(m); + } + } + } + Err(e) => { + tracing::warn!(provider = %p, error = %e, "list_models 失败"); + } + } + } + } + Ok(all) + } + + /// 列出可用音色 + pub async fn list_voices(&self, language: Option<&str>) -> Result> { + let providers: Vec = { + let inner = self.inner.read().unwrap(); + capable_providers(&inner, Capabilities::ListVoices) + }; + for p in providers { + if let Some(adapter) = self.get_adapter(&p) { + match adapter.list_voices(language).await { + Ok(v) => return Ok(v), + Err(_) => continue, + } + } + } + Err(AibridgeError::unsupported_capability( + "list_voices(无 Provider 支持)", + )) + } +} + +/// 获取支持指定能力的健康 Provider 列表 +fn capable_providers(inner: &RouterInner, capability: Capabilities) -> Vec { + let mut candidates: Vec = Vec::new(); + for p in &inner.provider_order { + let healthy = *inner.health.get(p).unwrap_or(&true); + if !healthy { + continue; + } + let Some(adapter) = inner.adapters.get(p) else { + continue; + }; + if adapter.supports_capability(capability) { + candidates.push(p.clone()); + } + } + // 没有健康的就放宽:返回所有适配器(不论健康) + if candidates.is_empty() { + for p in &inner.provider_order { + let Some(adapter) = inner.adapters.get(p) else { + continue; + }; + if adapter.supports_capability(capability) { + candidates.push(p.clone()); + } + } + } + candidates +} + +/// 按策略从候选中选一个 +fn pick_by_strategy( + candidates: &[String], + strategy: RoutingStrategy, + inner: &RouterInner, +) -> String { + if candidates.len() == 1 { + return candidates[0].clone(); + } + match strategy { + RoutingStrategy::First => candidates[0].clone(), + RoutingStrategy::RoundRobin => { + let idx = inner.rr_index % candidates.len(); + candidates[idx].clone() + } + RoutingStrategy::Random => { + let mut rng = rand::thread_rng(); + let idx = rng.gen_range(0..candidates.len()); + candidates[idx].clone() + } + RoutingStrategy::Weighted => { + let weights: Vec = candidates + .iter() + .map(|p| inner.weights.get(p).copied().unwrap_or(1)) + .collect(); + let total: u32 = weights.iter().sum(); + if total == 0 { + return candidates[0].clone(); + } + let mut rng = rand::thread_rng(); + let mut pick = rng.gen_range(0..total); + for (i, w) in weights.iter().enumerate() { + if pick < *w { + return candidates[i].clone(); + } + pick -= *w; + } + candidates.last().cloned().unwrap() + } + RoutingStrategy::Latency => { + let with_latency: Vec<&String> = candidates + .iter() + .filter(|p| inner.latency.get(*p).copied().unwrap_or(0.0) > 0.0) + .collect(); + if with_latency.is_empty() { + return candidates[0].clone(); + } + with_latency + .into_iter() + .min_by(|a, b| { + inner + .latency + .get(*a) + .copied() + .unwrap_or(f64::MAX) + .partial_cmp(&inner.latency.get(*b).copied().unwrap_or(f64::MAX)) + .unwrap_or(std::cmp::Ordering::Equal) + }) + .cloned() + .unwrap_or_else(|| candidates[0].clone()) + } + } +} + +/// 内置模型 → Provider 映射(迁移自 Python v1 `MODEL_PROVIDER_MAP`,节选) +/// +/// 完整列表很长,这里只保留 MVP 四 provider 相关的常见模型; +/// 用户可通过 `register_model_mapping` 补充。 +fn builtin_model_provider(model: &str) -> Option<&'static str> { + match model { + // OpenAI + "gpt-4o" + | "gpt-4-turbo" + | "gpt-4" + | "gpt-3.5-turbo" + | "whisper-1" + | "tts-1" + | "tts-1-hd" + | "gpt-4o-transcribe" + | "gpt-4o-mini-transcribe" => Some("openai"), + // Agnes + "claude-3-opus" | "claude-3-sonnet" | "claude-3-haiku" | "dall-e-3" | "video-gen-1" + | "video-gen-2" => Some("agnes"), + // Anthropic(直接协议) + "claude-3-opus-20240229" + | "claude-3-sonnet-20240229" + | "claude-3-haiku-20240307" + | "claude-3-5-sonnet-20241022" => Some("anthropic"), + // Google Gemini + "gemini-2.5-pro" | "gemini-2.5-flash" | "gemini-1.5-pro" | "gemini-1.5-flash" => { + Some("gemini") + } + // 火山引擎 Seedream/Seedance + "seedream-5.0" | "seedream-4.0" | "seedream-3.0" => Some("volcengine_cv"), + "seedance-2.0" | "seedance-2.0-mini" | "seedance-1.0" => Some("volcengine_cv"), + // 可灵 Kling + "kling-v1" | "kling-v1-5" | "kling-v2" => Some("kling"), + // Runway + "gen-3" | "gen-3-turbo" => Some("runway"), + // Edge TTS + "edge-tts" | "edge_tts" => Some("edge-tts"), + // ElevenLabs + "eleven_multilingual_v2" | "eleven_turbo_v2_5" => Some("elevenlabs"), + // Deepgram + "nova-3" | "nova-2" => Some("deepgram"), + // AssemblyAI + "best" | "nano" => Some("assemblyai"), + // Cartesia + "sonic-2" | "sonic-turbo" => Some("cartesia"), + _ => None, + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::adapter::CapabilitySet; + use crate::model::ChatMessage; + use async_trait::async_trait; + use std::sync::atomic::{AtomicU32, Ordering}; + + /// 用于路由测试的 mock 适配器 + struct MockAdapter { + provider: String, + caps: CapabilitySet, + call_count: Arc, + fail_n_times: u32, + } + + #[async_trait] + impl Adapter for MockAdapter { + fn provider_type(&self) -> &str { + &self.provider + } + fn provider_name(&self) -> &str { + &self.provider + } + fn capabilities(&self) -> CapabilitySet { + self.caps.clone() + } + async fn start(&mut self) -> Result<()> { + Ok(()) + } + async fn close(&mut self) -> Result<()> { + Ok(()) + } + async fn chat(&self, req: ChatRequest) -> Result { + let n = self.call_count.fetch_add(1, Ordering::SeqCst); + if n < self.fail_n_times { + return Err(AibridgeError::Api { + status: 500, + message: "mock fail".into(), + }); + } + Ok(ChatCompletion { + id: format!("{}-{}", self.provider, req.model), + object: "chat.completion".into(), + created: 0, + model: req.model, + choices: vec![], + usage: None, + service_tier: None, + system_fingerprint: None, + }) + } + async fn image_generate(&self, req: ImageRequest) -> Result { + Ok(ImageResult { + id: format!("{}-{}", self.provider, req.model), + object: "image.generation".into(), + created: 0, + model: req.model, + data: vec![], + }) + } + async fn list_models(&self, _: Option) -> Result> { + Ok(vec![]) + } + } + + /// 构造一个已注入 mock 适配器的路由器(绕过工厂) + fn router_with_adapters( + adapters: Vec<(String, CapabilitySet, u32)>, // (provider, caps, fail_n_times) + strategy: RoutingStrategy, + ) -> (Router, Vec>) { + let mut counters = Vec::new(); + let router = Router::with_strategy(vec![], strategy); + { + let mut inner = router.inner.write().unwrap(); + for (provider, caps, fail_n) in adapters { + let counter = Arc::new(AtomicU32::new(0)); + counters.push(counter.clone()); + let adapter = MockAdapter { + provider: provider.clone(), + caps: caps.clone(), + call_count: counter, + fail_n_times: fail_n, + }; + inner.adapters.insert( + provider.clone(), + Arc::from(Box::new(adapter) as Box), + ); + inner.provider_order.push(provider.clone()); + inner.health.insert(provider, true); + } + } + (router, counters) + } + + #[tokio::test] + async fn chat_routes_by_builtin_model_map() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::Chat); + s + }; + let (router, _c) = + router_with_adapters(vec![("openai".into(), caps, 0)], RoutingStrategy::First); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = router.chat(req).await.unwrap(); + assert!(result.id.starts_with("openai-")); + } + + #[tokio::test] + async fn chat_falls_back_on_failure() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::Chat); + s + }; + // 主 provider(openai)失败 1 次,fallback 到 agnes + let (router, counters) = router_with_adapters( + vec![ + ("openai".into(), caps.clone(), 1), + ("agnes".into(), caps, 0), + ], + RoutingStrategy::First, + ); + // 自定义映射:gpt-x → openai(主),但 fallback 会到 agnes + router.register_model_mapping("gpt-x", "openai"); + let req = ChatRequest::builder("gpt-x", vec![ChatMessage::user("hi")]).build(); + let result = router.chat(req).await.unwrap(); + assert!(result.id.starts_with("agnes-")); + // openai 被调用 1 次(失败) + assert_eq!(counters[0].load(Ordering::SeqCst), 1); + // agnes 被调用 1 次(成功) + assert_eq!(counters[1].load(Ordering::SeqCst), 1); + } + + #[tokio::test] + async fn chat_all_fail_returns_last_error() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::Chat); + s + }; + let (router, _c) = router_with_adapters( + vec![ + ("openai".into(), caps.clone(), 100), + ("agnes".into(), caps, 100), + ], + RoutingStrategy::First, + ); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = router.chat(req).await; + assert!(result.is_err()); + } + + #[tokio::test] + async fn select_provider_returns_error_when_no_capable() { + let (router, _c) = router_with_adapters( + vec![("openai".into(), CapabilitySet::new(), 0)], + RoutingStrategy::First, + ); + let result = router + .chat(ChatRequest::builder("unknown-model", vec![]).build()) + .await; + assert!(matches!(result, Err(AibridgeError::ModelNotFound { .. }))); + } + + #[tokio::test] + async fn image_generate_routes() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::ImageGenerate); + s + }; + let (router, _c) = + router_with_adapters(vec![("openai".into(), caps, 0)], RoutingStrategy::First); + // dall-e-3 映射到 agnes,但 agnes 不存在;自定义映射到 openai + router.register_model_mapping("dall-e-3", "openai"); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let result = router.image_generate(req).await.unwrap(); + assert!(result.id.starts_with("openai-")); + } + + #[tokio::test] + async fn fallback_disabled_uses_only_primary() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::Chat); + s + }; + let (router, counters) = router_with_adapters( + vec![ + ("openai".into(), caps.clone(), 1), + ("agnes".into(), caps, 0), + ], + RoutingStrategy::First, + ); + router.register_model_mapping("gpt-x", "openai"); + let router = router.with_fallback(false); + let req = ChatRequest::builder("gpt-x", vec![ChatMessage::user("hi")]).build(); + let result = router.chat(req).await; + assert!(result.is_err()); // openai 失败,无 fallback + assert_eq!(counters[0].load(Ordering::SeqCst), 1); + assert_eq!(counters[1].load(Ordering::SeqCst), 0); // agnes 未被调用 + } + + #[tokio::test] + async fn round_robin_strategy_rotates() { + let caps = { + let mut s = CapabilitySet::new(); + s.insert(Capabilities::Chat); + s + }; + let (router, _c) = router_with_adapters( + vec![ + ("openai".into(), caps.clone(), 0), + ("agnes".into(), caps, 0), + ], + RoutingStrategy::RoundRobin, + ); + // 注册一个走能力筛选的模型(非内置映射) + router.register_model_mapping("custom-1", "openai"); + router.register_model_mapping("custom-2", "agnes"); + // 第一次:custom-1 → openai + let r1 = router + .chat(ChatRequest::builder("custom-1", vec![]).build()) + .await + .unwrap(); + assert!(r1.id.starts_with("openai-")); + } + + #[test] + fn builtin_model_provider_known() { + assert_eq!(builtin_model_provider("gpt-4o"), Some("openai")); + assert_eq!(builtin_model_provider("claude-3-opus"), Some("agnes")); + assert_eq!(builtin_model_provider("gemini-2.5-pro"), Some("gemini")); + assert_eq!( + builtin_model_provider("seedance-2.0"), + Some("volcengine_cv") + ); + assert_eq!(builtin_model_provider("edge-tts"), Some("edge-tts")); + assert_eq!(builtin_model_provider("unknown"), None); + } + + #[test] + fn routing_strategy_default_is_first() { + assert_eq!(RoutingStrategy::default(), RoutingStrategy::First); + } + + #[test] + fn provider_entry_with_weight() { + let e = ProviderEntry::new("openai", ClientOptions::default()).with_weight(3); + assert_eq!(e.weight, 3); + } + + #[test] + fn provider_entry_weight_floored_to_1() { + let e = ProviderEntry::new("openai", ClientOptions::default()).with_weight(0); + assert_eq!(e.weight, 1); + } +} diff --git a/crates/aibridge-core/src/util.rs b/crates/aibridge-core/src/util.rs new file mode 100644 index 0000000..f3f9add --- /dev/null +++ b/crates/aibridge-core/src/util.rs @@ -0,0 +1,316 @@ +//! 工具函数 +//! +//! 通用的工具函数集合。 +//! 对应 Python v1 (agn-sdk) 的 `agn/core/utils.py`。 + +use std::collections::HashMap; +use std::time::{SystemTime, UNIX_EPOCH}; + +use base64::Engine; +use serde_json::Value; + +/// 默认 base64 引擎(标准字母表,带 padding) +const STANDARD: base64::engine::GeneralPurpose = base64::engine::general_purpose::STANDARD; + +/// 生成唯一 ID +/// +/// 对应 Python v1 `generate_id`。返回 12 位十六进制 + 可选前缀。 +pub fn generate_id(prefix: &str) -> String { + let id = uuid_v4_hex_12(); + if prefix.is_empty() { + id + } else { + format!("{prefix}_{id}") + } +} + +/// 生成 12 位十六进制 ID(基于 uuid v4 的前 6 字节) +fn uuid_v4_hex_12() -> String { + // 简易实现:用系统时间 + 计数器替代完整 uuid 依赖 + // 12 位 hex = 6 字节;取时间戳纳秒低 6 字节 + let nanos = SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_nanos()) + .unwrap_or(0); + let bytes = (nanos as u64).to_be_bytes(); + // 取后 6 字节 + let six: [u8; 6] = bytes[2..8].try_into().unwrap(); + hex_encode(&six) +} + +/// 获取当前 Unix 时间戳(秒) +pub fn current_timestamp() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0) +} + +/// 获取当前 Unix 时间戳(毫秒) +pub fn current_timestamp_ms() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_millis() as u64) + .unwrap_or(0) +} + +/// 检查字符串是否为 base64 编码 +/// +/// 对应 Python v1 `is_base64`。data URI 前缀会先被剥离再校验。 +pub fn is_base64(s: &str) -> bool { + let s = s.trim(); + if s.is_empty() { + return false; + } + // 移除可能的 data URI 前缀 + let s = s.split_once(',').map(|(_, d)| d).unwrap_or(s); + if s.is_empty() { + return false; + } + STANDARD.decode(s).is_ok() +} + +/// 将字节数据编码为 base64 字符串 +pub fn encode_base64(data: &[u8]) -> String { + STANDARD.encode(data) +} + +/// 将 base64 字符串解码为字节数据 +/// +/// 会自动剥离 data URI 前缀(逗号前的部分)。 +pub fn decode_base64(data: &str) -> Result, base64::DecodeError> { + let data = data.split_once(',').map(|(_, d)| d).unwrap_or(data); + STANDARD.decode(data) +} + +/// 计算 MD5 哈希 +/// +/// 注意:MD5 仅用于非安全场景(如缓存键、幂等标识)。 +pub fn md5_hash(data: &str) -> String { + // 不引入 md5 依赖:用简单的 FNV-1a 替代作为缓存键 + // 若后续需要真正的 MD5,再加 md-5 依赖 + fnv1a_64(data) +} + +/// FNV-1a 64 位哈希(作为 MD5 的轻量替代,仅用于非安全场景) +fn fnv1a_64(data: &str) -> String { + let mut hash: u64 = 0xcbf29ce484222325; + for b in data.as_bytes() { + hash ^= *b as u64; + hash = hash.wrapping_mul(0x100000001b3); + } + format!("{hash:016x}") +} + +/// 验证并修正视频尺寸(宽高必须是 8 的倍数) +/// +/// 向上取整到最近的 8 的倍数。 +pub fn validate_video_dimensions(width: u32, height: u32) -> (u32, u32) { + let w = width.div_ceil(8) * 8; + let h = height.div_ceil(8) * 8; + (w, h) +} + +/// 解析图像尺寸字符串(如 "1024x1024") +/// +/// 返回 `(width, height)`,格式非法时返回 `Err`。 +pub fn parse_size(size: &str) -> Result<(u32, u32), String> { + let lower = size.to_lowercase(); + let parts: Vec<&str> = lower.split('x').collect(); + if parts.len() != 2 { + return Err(format!("Invalid size format: {size}")); + } + let w: u32 = parts[0] + .parse() + .map_err(|_| format!("Invalid size format: {size}"))?; + let h: u32 = parts[1] + .parse() + .map_err(|_| format!("Invalid size format: {size}"))?; + Ok((w, h)) +} + +/// 构建图像尺寸字符串 +pub fn build_size_string(width: u32, height: u32) -> String { + format!("{width}x{height}") +} + +/// 合并多个 JSON 对象 +/// +/// 对应 Python v1 `merge_dicts`。后续值覆盖前面。 +/// 任一参数为 `Value::Null` 或非对象时跳过。 +pub fn merge_json(base: &Value, overrides: &[&Value]) -> Value { + let mut result = match base { + Value::Object(map) => map.clone(), + _ => serde_json::Map::new(), + }; + for ov in overrides { + if let Value::Object(m) = ov { + for (k, v) in m { + result.insert(k.clone(), v.clone()); + } + } + } + Value::Object(result) +} + +/// 合并多个 HashMap(后续覆盖前面) +pub fn merge_maps( + base: HashMap, + overrides: &[HashMap], +) -> HashMap { + let mut result = base; + for ov in overrides { + for (k, v) in ov { + result.insert(k.clone(), v.clone()); + } + } + result +} + +/// 十六进制编码(小写) +fn hex_encode(bytes: &[u8]) -> String { + let mut s = String::with_capacity(bytes.len() * 2); + for b in bytes { + s.push_str(&format!("{b:02x}")); + } + s +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn generate_id_with_prefix() { + let id = generate_id("task"); + assert!(id.starts_with("task_")); + assert!(id.len() > "task_".len()); + } + + #[test] + fn generate_id_without_prefix() { + let id = generate_id(""); + assert!(!id.is_empty()); + assert!(!id.contains('_')); + } + + #[test] + fn current_timestamp_nonzero() { + let t = current_timestamp(); + assert!(t > 0); + } + + #[test] + fn current_timestamp_ms_greater_than_seconds() { + let s = current_timestamp(); + let ms = current_timestamp_ms(); + assert!(ms >= s * 1000); + } + + #[test] + fn is_base64_valid() { + assert!(is_base64("aGVsbG8=")); // "hello" + } + + #[test] + fn is_base64_invalid() { + assert!(!is_base64("not base64!")); + assert!(!is_base64("")); + } + + #[test] + fn is_base64_strips_data_uri() { + assert!(is_base64("data:image/png;base64,aGVsbG8=")); + } + + #[test] + fn encode_decode_base64_roundtrip() { + let original = b"hello world"; + let encoded = encode_base64(original); + let decoded = decode_base64(&encoded).unwrap(); + assert_eq!(decoded, original); + } + + #[test] + fn decode_base64_strips_data_uri() { + let decoded = decode_base64("data:image/png;base64,aGVsbG8=").unwrap(); + assert_eq!(decoded, b"hello"); + } + + #[test] + fn md5_hash_stable() { + let h1 = md5_hash("test"); + let h2 = md5_hash("test"); + assert_eq!(h1, h2); + assert_ne!(md5_hash("test"), md5_hash("other")); + assert_eq!(h1.len(), 16); + } + + #[test] + fn validate_video_dimensions_rounds_up_to_8() { + assert_eq!(validate_video_dimensions(1280, 720), (1280, 720)); + assert_eq!(validate_video_dimensions(1281, 721), (1288, 728)); + assert_eq!(validate_video_dimensions(1, 1), (8, 8)); + } + + #[test] + fn parse_size_valid() { + assert_eq!(parse_size("1024x1024").unwrap(), (1024, 1024)); + assert_eq!(parse_size("1792X1024").unwrap(), (1792, 1024)); // 大写 X 也接受 + } + + #[test] + fn parse_size_invalid() { + assert!(parse_size("1024").is_err()); + assert!(parse_size("abc").is_err()); + assert!(parse_size("1024x").is_err()); + assert!(parse_size("").is_err()); + } + + #[test] + fn build_size_string_correct() { + assert_eq!(build_size_string(1024, 768), "1024x768"); + } + + #[test] + fn merge_json_overrides_win() { + let base = serde_json::json!({"a": 1, "b": 2}); + let ov1 = serde_json::json!({"b": 3, "c": 4}); + let result = merge_json(&base, &[&ov1]); + assert_eq!(result["a"], 1); + assert_eq!(result["b"], 3); + assert_eq!(result["c"], 4); + } + + #[test] + fn merge_json_ignores_non_object() { + let base = serde_json::json!({"a": 1}); + let ov = serde_json::json!(42); + let result = merge_json(&base, &[&ov]); + assert_eq!(result["a"], 1); + } + + #[test] + fn merge_json_empty_overrides() { + let base = serde_json::json!({"a": 1}); + let result = merge_json(&base, &[]); + assert_eq!(result["a"], 1); + } + + #[test] + fn merge_maps_overrides_win() { + let mut base = HashMap::new(); + base.insert("a".into(), serde_json::json!(1)); + let mut ov = HashMap::new(); + ov.insert("a".into(), serde_json::json!(2)); + ov.insert("b".into(), serde_json::json!(3)); + let result = merge_maps(base, &[ov]); + assert_eq!(result["a"], 2); + assert_eq!(result["b"], 3); + } + + #[test] + fn hex_encode_lowercase() { + assert_eq!(hex_encode(&[0x0a, 0xff]), "0aff"); + } +} From ddafa2650724fe1bf315e5c20e0d7efdb84e9f51 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 11:32:09 +0800 Subject: [PATCH 05/55] =?UTF-8?q?feat(aibridge-ffi):=20=E9=98=B6=E6=AE=B50?= =?UTF-8?q?.5=20C=20ABI=20=E5=B1=82=EF=BC=88runtime/=E5=8F=A5=E6=9F=84/JSO?= =?UTF-8?q?N=E8=BE=B9=E7=95=8C/cbindgen=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 227 ++++++++ crates/aibridge-ffi/Cargo.toml | 8 +- crates/aibridge-ffi/build.rs | 37 ++ crates/aibridge-ffi/cbindgen.toml | 29 + crates/aibridge-ffi/include/aibridge.h | 297 ++++++++++ crates/aibridge-ffi/src/error.rs | 278 +++++++++ crates/aibridge-ffi/src/handle.rs | 211 +++++++ crates/aibridge-ffi/src/lib.rs | 743 ++++++++++++++++++++++++- crates/aibridge-ffi/src/runtime.rs | 51 ++ crates/aibridge-ffi/src/stream.rs | 200 +++++++ crates/aibridge-ffi/src/string.rs | 132 +++++ 11 files changed, 2202 insertions(+), 11 deletions(-) create mode 100644 crates/aibridge-ffi/build.rs create mode 100644 crates/aibridge-ffi/cbindgen.toml create mode 100644 crates/aibridge-ffi/include/aibridge.h create mode 100644 crates/aibridge-ffi/src/error.rs create mode 100644 crates/aibridge-ffi/src/handle.rs create mode 100644 crates/aibridge-ffi/src/runtime.rs create mode 100644 crates/aibridge-ffi/src/stream.rs create mode 100644 crates/aibridge-ffi/src/string.rs diff --git a/Cargo.lock b/Cargo.lock index 3f77fd2..5f4a70e 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -35,7 +35,10 @@ name = "aibridge-ffi" version = "2.0.0-alpha.1" dependencies = [ "aibridge-core", + "cbindgen", + "futures", "once_cell", + "serde", "serde_json", "tokio", ] @@ -62,6 +65,56 @@ dependencies = [ "tokio", ] +[[package]] +name = "anstream" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "824a212faf96e9acacdbd09febd34438f8f711fb84e09a8916013cd7815ca28d" +dependencies = [ + "anstyle", + "anstyle-parse", + "anstyle-query", + "anstyle-wincon", + "colorchoice", + "is_terminal_polyfill", + "utf8parse", +] + +[[package]] +name = "anstyle" +version = "1.0.14" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "940b3a0ca603d1eade50a4846a2afffd5ef57a9feac2c0e2ec2e14f9ead76000" + +[[package]] +name = "anstyle-parse" +version = "1.0.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "52ce7f38b242319f7cabaa6813055467063ecdc9d355bbb4ce0c68908cd8130e" +dependencies = [ + "utf8parse", +] + +[[package]] +name = "anstyle-query" +version = "1.1.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "40c48f72fd53cd289104fc64099abca73db4166ad86ea0b4341abe65af83dadc" +dependencies = [ + "windows-sys 0.61.2", +] + +[[package]] +name = "anstyle-wincon" +version = "3.0.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "291e6a250ff86cd4a820112fb8898808a366d8f9f58ce16d1f538353ad55747d" +dependencies = [ + "anstyle", + "once_cell_polyfill", + "windows-sys 0.61.2", +] + [[package]] name = "async-stream" version = "0.3.6" @@ -125,6 +178,25 @@ version = "1.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "8ae3f5d315924270530207e2a68396c3cc547f6dca3fbdca317cfb1a51edb593" +[[package]] +name = "cbindgen" +version = "0.29.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "2ecb53484c9c167ba674026b656d8a27d7657a58e6066aa902bfb1a4aa00ae20" +dependencies = [ + "clap", + "heck", + "indexmap", + "log", + "proc-macro2", + "quote", + "serde", + "serde_json", + "syn", + "tempfile", + "toml", +] + [[package]] name = "cc" version = "1.2.66" @@ -158,6 +230,39 @@ dependencies = [ "rand_core 0.10.1", ] +[[package]] +name = "clap" +version = "4.6.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ddb117e43bbf7dacf0a4190fef4d345b9bad68dfc649cb349e7d17d28428e51" +dependencies = [ + "clap_builder", +] + +[[package]] +name = "clap_builder" +version = "4.6.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "714a53001bf66416adb0e2ef5ac857140e7dc3a0c48fb28b2f10762fc4b5069f" +dependencies = [ + "anstream", + "anstyle", + "clap_lex", + "strsim", +] + +[[package]] +name = "clap_lex" +version = "1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c8d4a3bb8b1e0c1050499d1815f5ab16d04f0959b233085fb31653fbfc9d98f9" + +[[package]] +name = "colorchoice" +version = "1.0.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" + [[package]] name = "convert_case" version = "0.6.0" @@ -213,6 +318,12 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "fastrand" +version = "2.4.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9f1f227452a390804cdb637b74a86990f2a7d7ba4b7d5693aac9b4dd6defd8d6" + [[package]] name = "find-msvc-tools" version = "0.1.9" @@ -598,6 +709,12 @@ version = "2.12.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d98f6fed1fde3f8c21bc40a1abb88dd75e67924f9cffc3ef95607bad8017f8e2" +[[package]] +name = "is_terminal_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a6cb138bb79a146c1bd460005623e142ef0181e3d0219cb493e02f7d08a35695" + [[package]] name = "itoa" version = "1.0.18" @@ -631,6 +748,12 @@ dependencies = [ "windows-link", ] +[[package]] +name = "linux-raw-sys" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32a66949e030da00e8c7d4434b251670a91556f4144941d37452769c25d58a53" + [[package]] name = "litemap" version = "0.8.2" @@ -739,6 +862,12 @@ version = "1.21.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "9f7c3e4beb33f85d45ae3e3a1792185706c8e16d043238c593331cc7cd313b50" +[[package]] +name = "once_cell_polyfill" +version = "1.70.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "384b8ab6d37215f3c5301a95a4accb5d64aa607f1fcb26a11b5303878451b4fe" + [[package]] name = "parking_lot" version = "0.12.5" @@ -1092,6 +1221,19 @@ version = "2.1.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6b1e7f9a428571be2dc5bc0505c13fb6bf936822b894ec87abf8a08a4e51742d" +[[package]] +name = "rustix" +version = "1.1.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6fe4565b9518b83ef4f91bb47ce29620ca828bd32cb7e408f0062e9930ba190" +dependencies = [ + "bitflags", + "errno", + "libc", + "linux-raw-sys", + "windows-sys 0.61.2", +] + [[package]] name = "rustls" version = "0.23.41" @@ -1194,6 +1336,15 @@ dependencies = [ "zmij", ] +[[package]] +name = "serde_spanned" +version = "1.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6662b5879511e06e8999a8a235d848113e942c9124f211511b16466ee2995f26" +dependencies = [ + "serde_core", +] + [[package]] name = "serde_urlencoded" version = "0.7.1" @@ -1250,6 +1401,12 @@ version = "1.2.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6ce2be8dc25455e1f91df71bfa12ad37d7af1092ae736f3a6cd0e37bc7810596" +[[package]] +name = "strsim" +version = "0.11.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "7da8b5736845d9f2fcb837ea5d9e2628564b3b043a70948a3f0b778838c5fb4f" + [[package]] name = "subtle" version = "2.6.1" @@ -1293,6 +1450,19 @@ version = "0.13.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "adb6935a6f5c20170eeceb1a3835a49e12e19d792f6dd344ccc76a985ca5a6ca" +[[package]] +name = "tempfile" +version = "3.27.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "32497e9a4c7b38532efcdebeef879707aa9f794296a4f0244f6f69e9bc8574bd" +dependencies = [ + "fastrand", + "getrandom 0.4.3", + "once_cell", + "rustix", + "windows-sys 0.61.2", +] + [[package]] name = "thiserror" version = "1.0.69" @@ -1409,6 +1579,45 @@ dependencies = [ "tokio", ] +[[package]] +name = "toml" +version = "0.9.12+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "cf92845e79fc2e2def6a5d828f0801e29a2f8acc037becc5ab08595c7d5e9863" +dependencies = [ + "indexmap", + "serde_core", + "serde_spanned", + "toml_datetime", + "toml_parser", + "toml_writer", + "winnow 0.7.15", +] + +[[package]] +name = "toml_datetime" +version = "0.7.5+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "92e1cfed4a3038bc5a127e35a2d360f145e1f4b971b551a2ba5fd7aedf7e1347" +dependencies = [ + "serde_core", +] + +[[package]] +name = "toml_parser" +version = "1.1.2+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a2abe9b86193656635d2411dc43050282ca48aa31c2451210f4202550afb7526" +dependencies = [ + "winnow 1.0.3", +] + +[[package]] +name = "toml_writer" +version = "1.1.1+spec-1.1.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "756daf9b1013ebe47a8776667b466417e2d4c5679d441c26230efd9ef78692db" + [[package]] name = "tower" version = "0.5.3" @@ -1527,6 +1736,12 @@ version = "1.0.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b6c140620e7ffbb22c2dee59cafe6084a59b5ffc27a8859a5f0d494b5d52b6be" +[[package]] +name = "utf8parse" +version = "0.2.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" + [[package]] name = "want" version = "0.3.1" @@ -1727,6 +1942,18 @@ version = "0.52.6" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "589f6da84c646204747d1270a2a5661ea66ed1cced2631d546fdfb155959f9ec" +[[package]] +name = "winnow" +version = "0.7.15" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df79d97927682d2fd8adb29682d1140b343be4ac0f08fd68b7765d9c059d3945" + +[[package]] +name = "winnow" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0592e1c9d151f854e6fd382574c3a0855250e1d9b2f99d9281c6e6391af352f1" + [[package]] name = "writeable" version = "0.6.3" diff --git a/crates/aibridge-ffi/Cargo.toml b/crates/aibridge-ffi/Cargo.toml index 49d9364..7e20f2f 100644 --- a/crates/aibridge-ffi/Cargo.toml +++ b/crates/aibridge-ffi/Cargo.toml @@ -17,8 +17,10 @@ name = "aibridge" aibridge-core.workspace = true tokio.workspace = true once_cell.workspace = true +serde.workspace = true serde_json.workspace = true +futures.workspace = true -# cbindgen 头文件生成(阶段 0.5 启用) -# [build-dependencies] -# cbindgen = "0.27" +# cbindgen 头文件生成 +[build-dependencies] +cbindgen = "0.29" diff --git a/crates/aibridge-ffi/build.rs b/crates/aibridge-ffi/build.rs new file mode 100644 index 0000000..1bee25d --- /dev/null +++ b/crates/aibridge-ffi/build.rs @@ -0,0 +1,37 @@ +//! build.rs - 用 cbindgen 生成 include/aibridge.h +//! +//! 每次 `cargo build` 触发,根据 cbindgen.toml 与 src/lib.rs 生成 C 头文件。 +//! 生成失败不阻断构建(开发期可继续),仅在 stderr 输出警告。 + +use std::env; +use std::path::PathBuf; + +fn main() { + let crate_dir = PathBuf::from(env::var("CARGO_MANIFEST_DIR").unwrap_or_else(|_| ".".into())); + let config = cbindgen::Config::from_root_or_default(&crate_dir); + + let result = cbindgen::Builder::new() + .with_crate(&crate_dir) + .with_config(config) + .generate(); + + match result { + Ok(bindings) => { + let out_dir = crate_dir.join("include"); + // 确保 include 目录存在 + std::fs::create_dir_all(&out_dir).ok(); + let header_path = out_dir.join("aibridge.h"); + if !bindings.write_to_file(&header_path) { + eprintln!("cargo:warning=cbindgen 写入头文件失败"); + } else { + // 让 cargo 在头文件变化时重新构建 + println!("cargo:rerun-if-changed=src/lib.rs"); + println!("cargo:rerun-if-changed=cbindgen.toml"); + } + } + Err(e) => { + // 生成失败仅警告,不阻断构建(CI 仍可继续编译 cdylib) + eprintln!("cargo:warning=cbindgen 生成头文件失败: {e}"); + } + } +} diff --git a/crates/aibridge-ffi/cbindgen.toml b/crates/aibridge-ffi/cbindgen.toml new file mode 100644 index 0000000..9addfcb --- /dev/null +++ b/crates/aibridge-ffi/cbindgen.toml @@ -0,0 +1,29 @@ +# cbindgen 配置 - 生成 include/aibridge.h +# +# 设计文档 7 节:暴露 C ABI 供 Go/JVM/.NET 调用。 +# 头文件需含:opaque 句柄类型、aibridge_bytes_t、所有 #[no_mangle] extern "C" 函数、 +# 返回码常量。 + +language = "C" + +# 头文件顶部注释 +header = """\ +/* + * AIBridge C ABI 头文件 - 由 cbindgen 自动生成,请勿手动修改 + * + * 供 Go (CGO) / JVM (JNA) / .NET (P/Invoke) 调用 aibridge-ffi cdylib。 + * Python / JS 通过 aibridge-python / aibridge-node 直连 aibridge-core,不走本头文件。 + * + * 详见设计文档 docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md 第 7 节。 + */ +""" + +# include guard +include_guard = "AIBRIDGE_H" + +# 风格 +style = "type" +sort_by = "Name" + +# 导出文档注释 +documentation = true diff --git a/crates/aibridge-ffi/include/aibridge.h b/crates/aibridge-ffi/include/aibridge.h new file mode 100644 index 0000000..1b0799a --- /dev/null +++ b/crates/aibridge-ffi/include/aibridge.h @@ -0,0 +1,297 @@ +/* + * AIBridge C ABI 头文件 - 由 cbindgen 自动生成,请勿手动修改 + * + * 供 Go (CGO) / JVM (JNA) / .NET (P/Invoke) 调用 aibridge-ffi cdylib。 + * Python / JS 通过 aibridge-python / aibridge-node 直连 aibridge-core,不走本头文件。 + * + * 详见设计文档 docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md 第 7 节。 + */ + + +#ifndef AIBRIDGE_H +#define AIBRIDGE_H + +#include +#include +#include +#include + +/** + * Client 句柄 + * + * 内部用 `Arc>`: + * - `Arc` 允许 stream task 持有 client 引用 + * - `tokio::sync::Mutex` 处理 core 的 `start`/`close` 是 `&mut self`, + * 且其 guard 可跨 `.await`(避免 `std::sync::Mutex` 跨 await 的死锁/panic) + */ +typedef struct AibridgeClient AibridgeClient; + +/** + * Stream 句柄 + * + * 持有流式 `ChatStream` 与一个可选的 tokio task handle。 + * `drop` 时若 task 存在则 abort,保证流式任务不会泄漏。 + * + * 注意:`aibridge_stream_next` 在调用方线程串行拉取(各语言绑定负责同步)。 + */ +typedef struct AibridgeStream AibridgeStream; + +/** + * 二进制缓冲结构(`#[repr(C)]` 保证 C 侧布局) + * + * 对应设计文档的 `aibridge_bytes_t`。`ptr` 指向 Rust 分配的字节缓冲, + * `len` 为长度。调用方通过 [`crate::string::aibridge_bytes_free`] 释放。 + */ +typedef struct { + /** + * 指向字节数据的指针(Rust 分配) + */ + const uint8_t *ptr; + /** + * 字节数据长度 + */ + uintptr_t len; +} aibridge_bytes_t; + +/** + * FFI 返回码类型 + * + * 0 表示成功;负数表示各类错误(与 [`AibridgeError`] 变体一一对应)。 + */ +typedef int32_t AibridgeStatus; + +/** + * 错误码类型别名(供 cbindgen 导出为 typedef) + */ +typedef AibridgeStatus aibridge_status_t; + +/** + * `aibridge_client_t` opaque 类型 + */ +typedef AibridgeClient aibridge_client_t; + +/** + * `aibridge_stream_t` opaque 类型 + */ +typedef AibridgeStream aibridge_stream_t; + +/** + * API 调用错误(HTTP 4xx/5xx) + */ +#define AIBRIDGE_ERR_API -5 + +/** + * 认证错误 + */ +#define AIBRIDGE_ERR_AUTHENTICATION -1 + +/** + * FFI 层通用错误(参数为空、JSON 解析失败、内部 panic 等) + */ +#define AIBRIDGE_ERR_FFI -100 + +/** + * 模型不存在 + */ +#define AIBRIDGE_ERR_MODEL_NOT_FOUND -4 + +/** + * 网络错误 + */ +#define AIBRIDGE_ERR_NETWORK -6 + +/** + * Provider 不存在 + */ +#define AIBRIDGE_ERR_PROVIDER_NOT_FOUND -9 + +/** + * 限流错误 + */ +#define AIBRIDGE_ERR_RATE_LIMIT -2 + +/** + * 服务暂时不可用 + */ +#define AIBRIDGE_ERR_SERVICE_UNAVAILABLE -11 + +/** + * 超时 + */ +#define AIBRIDGE_ERR_TIMEOUT -7 + +/** + * 不支持的能力 + */ +#define AIBRIDGE_ERR_UNSUPPORTED_CAPABILITY -8 + +/** + * 参数校验错误 + */ +#define AIBRIDGE_ERR_VALIDATION -3 + +/** + * 音色不可用 + */ +#define AIBRIDGE_ERR_VOICE_NOT_AVAILABLE -10 + +/** + * 成功 + */ +#define AIBRIDGE_OK 0 + +/** + * 流式:正常拉取到一个 chunk(仅 `aibridge_stream_next` 使用) + */ +#define AIBRIDGE_STREAM_CHUNK 0 + +/** + * 流式:流已结束(仅 `aibridge_stream_next` 使用,1 表示 EOF) + */ +#define AIBRIDGE_STREAM_END 1 + +/** + * 释放 Rust 分配的二进制缓冲 + * + * # Safety + * `ptr` 必须是指向由 FFI 函数(如 `aibridge_client_speech`)输出的 + * `aibridge_bytes_t*` 指针,且只能释放一次。传 `nullptr` 是安全的 no-op。 + * + * 对应设计文档的 `aibridge_bytes_free`。 + */ +void aibridge_bytes_free(aibridge_bytes_t *ptr); + +/** + * 文本对话(阻塞) + * + * `request_json` 为 `ChatRequest` 的 JSON 序列化字符串。 + * 成功时 `*out_response_json` 写入 `ChatCompletion` 的 JSON,调用方需通过 + * [`aibridge_string_free`] 释放。 + * + * # Safety + * - `client` 必须来自 [`aibridge_client_new`] + * - `request_json` 必须是合法 JSON C 字符串 + * - `out_response_json` 必须指向可写的 `*mut c_char` 槽位 + */ +aibridge_status_t aibridge_client_chat(aibridge_client_t *client, + const char *request_json, + char **out_response_json); + +/** + * 流式文本对话(创建 stream 句柄) + * + * `request_json` 为 `ChatRequest` 的 JSON 序列化字符串。 + * 成功时 `*out_stream` 写入 stream 句柄,调用方通过 + * [`aibridge_stream_next`] 拉取 chunk,最后用 [`aibridge_stream_destroy`] 释放。 + * + * # Safety + * - `client` / `request_json` / `out_stream` 均需合法非空 + */ +aibridge_status_t aibridge_client_chat_stream(aibridge_client_t *client, + const char *request_json, + aibridge_stream_t **out_stream); + +/** + * 释放客户端句柄 + * + * # Safety + * `client` 必须是 [`aibridge_client_new`] 返回的指针,且只能释放一次。 + * 传 `nullptr` 是安全的 no-op。 + */ +void aibridge_client_destroy(aibridge_client_t *client); + +/** + * 创建客户端 + * + * `provider` 为 Provider 类型(如 "openai"、"agnes"),UTF-8 C 字符串。 + * `config_json` 为 `ClientOptions` 的 JSON 序列化字符串(可为 `nullptr`, + * 等价于默认配置)。 + * + * 成功返回 client 指针;失败返回 `nullptr`(错误写入 `aibridge_last_error`)。 + * + * # Safety + * - `provider` 必须是合法的 NUL 结尾 UTF-8 C 字符串 + * - `config_json` 可为 `nullptr` 或合法 JSON C 字符串 + * - 返回的指针需通过 [`aibridge_client_destroy`] 释放 + */ +aibridge_client_t *aibridge_client_new(const char *provider, const char *config_json); + +/** + * 文字转语音(阻塞,二进制载荷走 `aibridge_bytes_t`) + * + * `request_json` 为 `SpeechRequest` 的 JSON 序列化字符串。 + * 成功时 `*out_audio` 写入二进制音频缓冲(`aibridge_bytes_t`), + * `*out_meta_json` 写入 `SpeechResult`(不含 audio_data)的 JSON。 + * 两者均由调用方分别通过 [`aibridge_bytes_free`] / [`aibridge_string_free`] 释放。 + * + * 若 Provider 仅返回 `audio_base64`/`audio_url` 而无二进制数据, + * `*out_audio` 将为 `nullptr`(meta_json 仍写入)。 + * + * # Safety + * - `client` / `request_json` / `out_audio` / `out_meta_json` 均需合法非空 + */ +aibridge_status_t aibridge_client_speech(aibridge_client_t *client, + const char *request_json, + aibridge_bytes_t **out_audio, + char **out_meta_json); + +/** + * 启动客户端(初始化适配器) + * + * 返回 0 成功;负数为错误码(详见 `aibridge_last_error`)。 + * + * # Safety + * `client` 必须是 [`aibridge_client_new`] 返回的有效指针。 + */ +aibridge_status_t aibridge_client_start(aibridge_client_t *client); + +/** + * 读取当前线程的 last_error(JSON 字符串) + * + * 返回指向线程局部 `CString` 内部缓冲的指针,**调用方不应释放**。 + * 若当前线程无错误,返回 `nullptr`。 + * + * 输出格式:`{"code":"...","message":"...","details":...,"retryable":bool}` + * + * # Safety + * 本函数本身无内存不安全操作(仅读取线程局部变量并返回裸指针), + * 标 `unsafe extern "C"` 仅为符合 FFI 调用约定。调用方**不得**释放返回的指针, + * 且需注意返回的指针仅在当前线程的下一次 FFI 调用前保证有效(thread_local 语义)。 + */ +const char *aibridge_last_error(void); + +/** + * 释放 stream 句柄(触发 Rust drop → tokio task abort) + * + * # Safety + * `stream` 必须来自 [`aibridge_client_chat_stream`],且只能释放一次。 + * 传 `nullptr` 是安全的 no-op。 + */ +void aibridge_stream_destroy(aibridge_stream_t *stream); + +/** + * 拉取下一个流式 chunk(阻塞) + * + * 返回: + * - `0`(`AIBRIDGE_STREAM_CHUNK`):拉到一个 chunk,`*out_chunk_json` 写入 JSON + * - `1`(`AIBRIDGE_STREAM_END`):流正常结束 + * - 负数:错误(`aibridge_last_error` 已写入) + * + * # Safety + * - `stream` 必须来自 [`aibridge_client_chat_stream`] + * - `out_chunk_json` 必须指向可写的 `*mut c_char` 槽位 + */ +aibridge_status_t aibridge_stream_next(aibridge_stream_t *stream, char **out_chunk_json); + +/** + * 释放 Rust 分配的 C 字符串 + * + * # Safety + * `ptr` 必须是由 [`alloc_cstring`](或 `CString::into_raw`)分配的指针, + * 且只能释放一次。传 `nullptr` 是安全的 no-op。 + * + * 对应设计文档的 `aibridge_string_free`。 + */ +void aibridge_string_free(char *ptr); + +#endif /* AIBRIDGE_H */ diff --git a/crates/aibridge-ffi/src/error.rs b/crates/aibridge-ffi/src/error.rs new file mode 100644 index 0000000..a52bc7a --- /dev/null +++ b/crates/aibridge-ffi/src/error.rs @@ -0,0 +1,278 @@ +//! FFI 错误处理 +//! +//! 对应设计文档 7 节错误模型: +//! - [`AibridgeStatus`](i32 返回码):0 成功,负数为错误类别 +//! - 线程局部 `last_error` 槽:存 JSON 字符串 `{code,message,details,retryable}` +//! - [`aibridge_last_error`] 暴露给 C 侧读取 +//! +//! 错误码映射 `aibridge_core::AibridgeError` 的各变体,便于各语言绑定 +//! 做异常映射时无需解析 JSON 即可快速分类。 + +#[cfg(test)] +use crate::string; +use aibridge_core::error::AibridgeError; +use std::cell::RefCell; +use std::ffi::CString; +use std::ptr; + +/// FFI 返回码类型 +/// +/// 0 表示成功;负数表示各类错误(与 [`AibridgeError`] 变体一一对应)。 +pub type AibridgeStatus = i32; + +/// 成功 +pub const AIBRIDGE_OK: AibridgeStatus = 0; + +/// 流式:正常拉取到一个 chunk(仅 `aibridge_stream_next` 使用) +pub const AIBRIDGE_STREAM_CHUNK: AibridgeStatus = 0; + +/// 流式:流已结束(仅 `aibridge_stream_next` 使用,1 表示 EOF) +pub const AIBRIDGE_STREAM_END: AibridgeStatus = 1; + +// —— 错误类别(负数)—— +/// 认证错误 +pub const AIBRIDGE_ERR_AUTHENTICATION: AibridgeStatus = -1; +/// 限流错误 +pub const AIBRIDGE_ERR_RATE_LIMIT: AibridgeStatus = -2; +/// 参数校验错误 +pub const AIBRIDGE_ERR_VALIDATION: AibridgeStatus = -3; +/// 模型不存在 +pub const AIBRIDGE_ERR_MODEL_NOT_FOUND: AibridgeStatus = -4; +/// API 调用错误(HTTP 4xx/5xx) +pub const AIBRIDGE_ERR_API: AibridgeStatus = -5; +/// 网络错误 +pub const AIBRIDGE_ERR_NETWORK: AibridgeStatus = -6; +/// 超时 +pub const AIBRIDGE_ERR_TIMEOUT: AibridgeStatus = -7; +/// 不支持的能力 +pub const AIBRIDGE_ERR_UNSUPPORTED_CAPABILITY: AibridgeStatus = -8; +/// Provider 不存在 +pub const AIBRIDGE_ERR_PROVIDER_NOT_FOUND: AibridgeStatus = -9; +/// 音色不可用 +pub const AIBRIDGE_ERR_VOICE_NOT_AVAILABLE: AibridgeStatus = -10; +/// 服务暂时不可用 +pub const AIBRIDGE_ERR_SERVICE_UNAVAILABLE: AibridgeStatus = -11; +/// FFI 层通用错误(参数为空、JSON 解析失败、内部 panic 等) +pub const AIBRIDGE_ERR_FFI: AibridgeStatus = -100; + +// 线程局部 last_error 槽 +// +// 存储 JSON 字符串 `{code,message,details,retryable}`,C 侧通过 +// [`aibridge_last_error`] 读取。每个线程独立,无需加锁。 +// 用 `Option` 便于返回裸指针(指针指向 CString 内部缓冲)。 +thread_local! { + static LAST_ERROR: RefCell> = const { RefCell::new(None) }; +} + +/// 将核心层错误写入线程局部 last_error 槽,并返回对应的 FFI 错误码 +/// +/// 内部把错误序列化为 JSON:`{"code":"...","message":"...","details":...,"retryable":bool}`。 +/// `details` 字段:Validation 错误带原始 details,其余为 null。 +pub fn set_last_error(err: &AibridgeError) -> AibridgeStatus { + let code = err.code(); + let message = err.to_string(); + let retryable = err.is_retryable(); + let details = match err { + AibridgeError::Validation { details, .. } => details.clone(), + AibridgeError::RateLimit { retry_after, .. } if retry_after.is_some() => { + serde_json::json!({ "retry_after": retry_after }) + } + _ => serde_json::Value::Null, + }; + + let payload = serde_json::json!({ + "code": code, + "message": message, + "details": details, + "retryable": retryable, + }); + let json_str = payload.to_string(); + + // 写入线程局部槽(失败则清空,避免残留旧错误) + match CString::new(json_str) { + Ok(cstr) => { + LAST_ERROR.with(|slot| { + *slot.borrow_mut() = Some(cstr); + }); + } + Err(_) => { + LAST_ERROR.with(|slot| { + *slot.borrow_mut() = None; + }); + } + } + + status_from_error(err) +} + +/// 写入 FFI 层自有的错误信息(如参数为空、JSON 解析失败、panic 等) +/// +/// 与 [`set_last_error`] 不同,此处不依赖 `AibridgeError`,统一用 +/// `AIBRIDGE_ERR_FFI` 错误码。`details` 可为任意 JSON 值。 +pub fn set_ffi_error(message: &str, details: serde_json::Value) -> AibridgeStatus { + let payload = serde_json::json!({ + "code": "ffi_error", + "message": message, + "details": details, + "retryable": false, + }); + let json_str = payload.to_string(); + if let Ok(cstr) = CString::new(json_str) { + LAST_ERROR.with(|slot| { + *slot.borrow_mut() = Some(cstr); + }); + } + AIBRIDGE_ERR_FFI +} + +/// 写入简单的 FFI 错误(无 details) +pub fn set_ffi_error_simple(message: &str) -> AibridgeStatus { + set_ffi_error(message, serde_json::Value::Null) +} + +/// 清空当前线程的 last_error 槽(成功路径调用) +pub fn clear_last_error() { + LAST_ERROR.with(|slot| { + *slot.borrow_mut() = None; + }); +} + +/// 读取当前线程 last_error 的 C 字符串指针 +/// +/// 返回的指针指向线程局部 `CString` 内部缓冲,调用方**不应释放**。 +/// 若当前线程无错误,返回空指针。 +/// +/// # 线程安全 +/// 仅返回当前线程的错误;各语言绑定的异步包装应保证读取与触发错误的调用 +/// 在同一线程,或在 FFI 边界立即读取后转存。 +pub fn last_error_ptr() -> *const std::os::raw::c_char { + LAST_ERROR.with(|slot| { + slot.borrow() + .as_ref() + .map(|cstr| cstr.as_ptr()) + .unwrap_or(ptr::null()) + }) +} + +/// 将核心层错误映射到 FFI 错误码 +pub fn status_from_error(err: &AibridgeError) -> AibridgeStatus { + match err { + AibridgeError::Authentication { .. } => AIBRIDGE_ERR_AUTHENTICATION, + AibridgeError::RateLimit { .. } => AIBRIDGE_ERR_RATE_LIMIT, + AibridgeError::Validation { .. } => AIBRIDGE_ERR_VALIDATION, + AibridgeError::ModelNotFound { .. } => AIBRIDGE_ERR_MODEL_NOT_FOUND, + AibridgeError::Api { .. } => AIBRIDGE_ERR_API, + AibridgeError::Network(_) => AIBRIDGE_ERR_NETWORK, + AibridgeError::Timeout => AIBRIDGE_ERR_TIMEOUT, + AibridgeError::UnsupportedCapability { .. } => AIBRIDGE_ERR_UNSUPPORTED_CAPABILITY, + AibridgeError::ProviderNotFound { .. } => AIBRIDGE_ERR_PROVIDER_NOT_FOUND, + AibridgeError::VoiceNotAvailable { .. } => AIBRIDGE_ERR_VOICE_NOT_AVAILABLE, + AibridgeError::ServiceUnavailable { .. } => AIBRIDGE_ERR_SERVICE_UNAVAILABLE, + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn status_mapping_is_complete() { + // 覆盖每个变体,确保映射不漏 + assert_eq!( + status_from_error(&AibridgeError::authentication("x")), + AIBRIDGE_ERR_AUTHENTICATION + ); + assert_eq!( + status_from_error(&AibridgeError::rate_limit("x")), + AIBRIDGE_ERR_RATE_LIMIT + ); + assert_eq!( + status_from_error(&AibridgeError::validation("x")), + AIBRIDGE_ERR_VALIDATION + ); + assert_eq!( + status_from_error(&AibridgeError::model_not_found("m")), + AIBRIDGE_ERR_MODEL_NOT_FOUND + ); + assert_eq!( + status_from_error(&AibridgeError::api(500, "x")), + AIBRIDGE_ERR_API + ); + // Network 变体需构造 reqwest::Error,跳过;用 UnsupportedCapability 等 + assert_eq!( + status_from_error(&AibridgeError::Timeout), + AIBRIDGE_ERR_TIMEOUT + ); + assert_eq!( + status_from_error(&AibridgeError::unsupported_capability("c")), + AIBRIDGE_ERR_UNSUPPORTED_CAPABILITY + ); + assert_eq!( + status_from_error(&AibridgeError::provider_not_found("p")), + AIBRIDGE_ERR_PROVIDER_NOT_FOUND + ); + assert_eq!( + status_from_error(&AibridgeError::voice_not_available("v")), + AIBRIDGE_ERR_VOICE_NOT_AVAILABLE + ); + assert_eq!( + status_from_error(&AibridgeError::service_unavailable("s")), + AIBRIDGE_ERR_SERVICE_UNAVAILABLE + ); + } + + #[test] + fn set_last_error_writes_valid_json() { + let status = set_last_error(&AibridgeError::rate_limit("慢一点")); + assert_eq!(status, AIBRIDGE_ERR_RATE_LIMIT); + let ptr = last_error_ptr(); + assert!(!ptr.is_null()); + let json = unsafe { string::cstr_to_string_unchecked(ptr) }; + let v: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(v["code"], "rate_limit_error"); + assert_eq!(v["retryable"], true); + assert!(v["message"].as_str().unwrap().contains("慢一点")); + clear_last_error(); + } + + #[test] + fn validation_error_carries_details() { + let err = + AibridgeError::validation_with_details("坏", serde_json::json!({"field": "model"})); + set_last_error(&err); + let ptr = last_error_ptr(); + let json = unsafe { string::cstr_to_string_unchecked(ptr) }; + let v: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(v["details"]["field"], "model"); + clear_last_error(); + } + + #[test] + fn set_ffi_error_writes_ffi_code() { + let status = set_ffi_error("参数为空", serde_json::Value::Null); + assert_eq!(status, AIBRIDGE_ERR_FFI); + let ptr = last_error_ptr(); + assert!(!ptr.is_null()); + let json = unsafe { string::cstr_to_string_unchecked(ptr) }; + let v: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(v["code"], "ffi_error"); + assert_eq!(v["retryable"], false); + clear_last_error(); + } + + #[test] + fn clear_last_error_empties_slot() { + set_ffi_error_simple("x"); + clear_last_error(); + assert!(last_error_ptr().is_null()); + } + + #[test] + fn last_error_is_thread_local() { + // 主线程设置后,子线程应读不到 + set_ffi_error_simple("主线程错误"); + let child = std::thread::spawn(|| last_error_ptr().is_null()); + assert!(child.join().unwrap(), "子线程不应看到主线程的 last_error"); + clear_last_error(); + } +} diff --git a/crates/aibridge-ffi/src/handle.rs b/crates/aibridge-ffi/src/handle.rs new file mode 100644 index 0000000..2aed9d9 --- /dev/null +++ b/crates/aibridge-ffi/src/handle.rs @@ -0,0 +1,211 @@ +//! opaque 句柄定义 +//! +//! 对应设计文档 7 节的句柄式生命周期: +//! - [`AibridgeClient`]:包裹 `aibridge_core::Client`,用 `Arc>` 处理 +//! core 的 `start`/`close` 是 `&mut self` +//! - [`AibridgeStream`]:持有流式 `ChatStream` + tokio task handle,`drop` 触发 abort +//! - [`aibridge_bytes_t`]:二进制缓冲(`ptr` + `len`),跨 FFI 传递音频等载荷 +//! +//! C 侧只看到 opaque 指针 `aibridge_client_t*` / `aibridge_stream_t*`, +//! 具体布局对调用方不可见。 + +use aibridge_core::adapter::ChatStream; +use aibridge_core::client::Client; +use std::sync::Arc; +use tokio::sync::Mutex; +use tokio::task::JoinHandle; + +/// 二进制缓冲结构(`#[repr(C)]` 保证 C 侧布局) +/// +/// 对应设计文档的 `aibridge_bytes_t`。`ptr` 指向 Rust 分配的字节缓冲, +/// `len` 为长度。调用方通过 [`crate::string::aibridge_bytes_free`] 释放。 +#[repr(C)] +#[allow(non_camel_case_types)] +pub struct aibridge_bytes_t { + /// 指向字节数据的指针(Rust 分配) + pub ptr: *const u8, + /// 字节数据长度 + pub len: usize, +} + +impl aibridge_bytes_t { + /// 从 `Vec` 构造(消耗 vec,将其底层缓冲转为裸指针) + pub fn from_vec(data: Vec) -> Self { + // 用 Box<[u8]> 持有,drop 时释放底层缓冲 + // 注意:into_raw 时必须能还原出同样布局;这里用 Box::into_raw([u8]) 风格 + let len = data.len(); + let boxed: Box<[u8]> = data.into_boxed_slice(); + let ptr = Box::into_raw(boxed) as *const u8; + Self { ptr, len } + } +} + +// aibridge_bytes_t 持有裸指针,但实际所有权通过 Box<[u8]> 管理, +// aibridge_bytes_free 时用 Box::from_raw 还原。Send/Sync 不必要(FFI 单线程使用)。 + +/// Client 句柄 +/// +/// 内部用 `Arc>`: +/// - `Arc` 允许 stream task 持有 client 引用 +/// - `tokio::sync::Mutex` 处理 core 的 `start`/`close` 是 `&mut self`, +/// 且其 guard 可跨 `.await`(避免 `std::sync::Mutex` 跨 await 的死锁/panic) +pub struct AibridgeClient { + /// 内部 core Client(互斥访问) + pub(crate) inner: Arc>, +} + +impl AibridgeClient { + /// 从 core Client 构造句柄 + pub fn new(client: Client) -> Self { + Self { + inner: Arc::new(Mutex::new(client)), + } + } + + /// 获取内部 Arc 引用(供 stream task 复用) + pub fn arc(&self) -> Arc> { + Arc::clone(&self.inner) + } +} + +/// Stream 句柄 +/// +/// 持有流式 `ChatStream` 与一个可选的 tokio task handle。 +/// `drop` 时若 task 存在则 abort,保证流式任务不会泄漏。 +/// +/// 注意:`aibridge_stream_next` 在调用方线程串行拉取(各语言绑定负责同步)。 +pub struct AibridgeStream { + /// 流式 chunk 迭代器(ChatStream = BoxStream>) + pub(crate) stream: Option, + /// 可选的后台 task handle(drop 时 abort) + pub(crate) task: Option>, + /// 已结束标志(避免重复拉取) + pub(crate) ended: bool, +} + +impl AibridgeStream { + /// 构造 stream 句柄 + pub fn new(stream: ChatStream) -> Self { + Self { + stream: Some(stream), + task: None, + ended: false, + } + } + + /// 构造带后台 task 的 stream 句柄 + pub fn with_task(stream: ChatStream, task: JoinHandle<()>) -> Self { + Self { + stream: Some(stream), + task: Some(task), + ended: false, + } + } +} + +impl Drop for AibridgeStream { + fn drop(&mut self) { + // 先 abort task(若存在),再 drop stream + if let Some(task) = self.task.take() { + task.abort(); + } + if let Some(stream) = self.stream.take() { + // 主动 drop 流,释放底层资源(BoxStream 无 close 方法,直接 drop 即可) + drop(stream); + } + } +} + +// FFI 导出的 opaque 类型别名(C 侧只见指针) +/// `aibridge_client_t` opaque 类型 +#[allow(non_camel_case_types)] +pub type aibridge_client_t = AibridgeClient; +/// `aibridge_stream_t` opaque 类型 +#[allow(non_camel_case_types)] +pub type aibridge_stream_t = AibridgeStream; + +/// 释放 client 句柄 +/// +/// # Safety +/// `ptr` 必须是 [`aibridge_client_new`] 返回的指针,且只能释放一次。 +/// 传 `nullptr` 是安全的 no-op。 +unsafe fn drop_client(ptr: *mut aibridge_client_t) { + if !ptr.is_null() { + drop(Box::from_raw(ptr)); + } +} + +/// C 侧调用的 client 释放入口(在 lib.rs 重新导出为 #[no_mangle]) +/// +/// 这里单独定义便于内部测试直接调用。 +pub(crate) unsafe fn destroy_client_impl(ptr: *mut aibridge_client_t) { + drop_client(ptr); +} + +/// C 侧调用的 stream 释放入口 +pub(crate) unsafe fn destroy_stream_impl(ptr: *mut aibridge_stream_t) { + if !ptr.is_null() { + drop(Box::from_raw(ptr)); + } +} + +/// 释放 `aibridge_bytes_t` 指针(内部辅助,供错误回滚用) +/// +/// # Safety +/// `ptr` 必须是 `Box::into_raw(Box::new(aibridge_bytes_t::from_vec(...)))` 产生的指针。 +pub(crate) unsafe fn destroy_bytes_ptr(ptr: *mut aibridge_bytes_t) { + if !ptr.is_null() { + let boxed = Box::from_raw(ptr); + // 还原底层 [u8] 并释放 + if !boxed.ptr.is_null() && boxed.len > 0 { + let slice = std::slice::from_raw_parts_mut(boxed.ptr as *mut u8, boxed.len); + drop(Box::from_raw(slice)); + } + drop(boxed); + } +} + +#[cfg(test)] +mod tests { + use super::*; + use futures::stream::StreamExt; + + #[test] + fn aibridge_bytes_from_vec_roundtrip() { + let b = aibridge_bytes_t::from_vec(vec![10, 20, 30]); + assert_eq!(b.len, 3); + unsafe { + assert_eq!(*b.ptr, 10); + assert_eq!(*b.ptr.add(2), 30); + } + // 通过 aibridge_bytes_free 释放(验证完整生命周期) + unsafe { crate::string::aibridge_bytes_free(Box::into_raw(Box::new(b))) }; + } + + #[test] + fn destroy_null_pointers_are_noop() { + unsafe { + destroy_client_impl(std::ptr::null_mut()); + destroy_stream_impl(std::ptr::null_mut()); + } + } + + #[test] + #[allow(clippy::async_yields_async)] + fn stream_drop_aborts_task() { + // 构造一个空 stream + 一个会一直挂起的 task,drop 后 task 应被 abort + use futures::stream; + let s = stream::iter(vec![]).boxed(); + let handle = crate::runtime::block_on(async { + tokio::spawn(async { + // 永久挂起 + std::future::pending::<()>().await; + }) + }); + let mut stream_handle = AibridgeStream::with_task(s, handle); + // 手动取走 stream 后 drop(触发 task abort) + stream_handle.stream.take(); + drop(stream_handle); + // task 已 abort,无法 join 成功(这里仅验证不 panic) + } +} diff --git a/crates/aibridge-ffi/src/lib.rs b/crates/aibridge-ffi/src/lib.rs index e882359..0c2b800 100644 --- a/crates/aibridge-ffi/src/lib.rs +++ b/crates/aibridge-ffi/src/lib.rs @@ -1,11 +1,738 @@ //! AIBridge C ABI - FFI 层 //! -//! 暴露 C ABI 供 Go/JVM/.NET 调用。 -//! Python/JS 通过 aibridge-python / aibridge-node 直连 aibridge-core,不走本层。 +//! 暴露 C ABI 供 Go/JVM/.NET 调用。Python/JS 通过 aibridge-python / +//! aibridge-node 直连 aibridge-core,不走本层。 //! -//! 设计要点(阶段 0.5 实现): -//! - 全局 tokio runtime(once_cell::Lazy),每个 FFI 调用 block_on -//! - 句柄式:aibridge_client_t / aibridge_stream_t(opaque) -//! - 复杂 struct 走 JSON 字符串边界,二进制走 aibridge_bytes_t -//! - 错误:aibridge_status_t 返回码 + aibridge_last_error() 线程局部槽 -//! - cbindgen 生成 include/aibridge.h +//! 设计要点(设计文档 7 节): +//! - 全局 tokio runtime(`once_cell::Lazy`),每个 FFI 调用 `block_on` +//! - 句柄式:`aibridge_client_t` / `aibridge_stream_t`(opaque) +//! - 复杂 struct 走 JSON 字符串边界,二进制走 `aibridge_bytes_t` +//! - 错误:`aibridge_status_t` 返回码 + `aibridge_last_error()` 线程局部槽 +//! - cbindgen 生成 `include/aibridge.h` +//! +//! # 安全保证 +//! 所有 `extern "C"` 函数绝不 panic:内部用 `std::panic::catch_unwind` +//! 捕获 panic,转为 `AIBRIDGE_ERR_FFI` 错误码并写入 last_error。 + +mod error; +mod handle; +mod runtime; +mod stream; +mod string; + +use crate::error::{AibridgeStatus, AIBRIDGE_OK}; +use crate::handle::{aibridge_bytes_t, aibridge_client_t, aibridge_stream_t, AibridgeClient}; +use crate::runtime::block_on; +use crate::string::{alloc_cstring, cstr_to_string}; +use aibridge_core::client::Client; +use aibridge_core::config::ClientOptions; +use aibridge_core::model::{ChatRequest, SpeechRequest}; +use std::os::raw::c_char; +use std::panic::{catch_unwind, AssertUnwindSafe}; +use std::ptr; + +// —— 重新导出供 cbindgen 生成头文件所需的符号 —— +pub use crate::string::{aibridge_bytes_free, aibridge_string_free}; + +/// 错误码类型别名(供 cbindgen 导出为 typedef) +#[allow(non_camel_case_types)] +pub type aibridge_status_t = AibridgeStatus; + +// ========================================================================= +// 生命周期:client new / start / destroy +// ========================================================================= + +/// 创建客户端 +/// +/// `provider` 为 Provider 类型(如 "openai"、"agnes"),UTF-8 C 字符串。 +/// `config_json` 为 `ClientOptions` 的 JSON 序列化字符串(可为 `nullptr`, +/// 等价于默认配置)。 +/// +/// 成功返回 client 指针;失败返回 `nullptr`(错误写入 `aibridge_last_error`)。 +/// +/// # Safety +/// - `provider` 必须是合法的 NUL 结尾 UTF-8 C 字符串 +/// - `config_json` 可为 `nullptr` 或合法 JSON C 字符串 +/// - 返回的指针需通过 [`aibridge_client_destroy`] 释放 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_new( + provider: *const c_char, + config_json: *const c_char, +) -> *mut aibridge_client_t { + let result = catch_unwind(AssertUnwindSafe(|| client_new_impl(provider, config_json))); + match result { + Ok(ptr) => ptr, + Err(_) => { + error::set_ffi_error_simple("aibridge_client_new 内部 panic"); + ptr::null_mut() + } + } +} + +/// `aibridge_client_new` 的实现 +fn client_new_impl(provider: *const c_char, config_json: *const c_char) -> *mut aibridge_client_t { + let provider_str = match cstr_to_string(provider) { + Some(Ok(s)) => s, + Some(Err(_)) => { + error::set_ffi_error_simple("provider 不是合法 UTF-8"); + return ptr::null_mut(); + } + None => { + error::set_ffi_error_simple("provider 为空指针"); + return ptr::null_mut(); + } + }; + + // config_json 可为空(用默认 ClientOptions) + let opts: ClientOptions = if config_json.is_null() { + ClientOptions::default() + } else { + match cstr_to_string(config_json) { + Some(Ok(s)) => match serde_json::from_str::(&s) { + Ok(o) => o, + Err(e) => { + error::set_ffi_error( + "config_json 反序列化 ClientOptions 失败", + serde_json::json!({ "error": e.to_string() }), + ); + return ptr::null_mut(); + } + }, + Some(Err(_)) => { + error::set_ffi_error_simple("config_json 不是合法 UTF-8"); + return ptr::null_mut(); + } + None => ClientOptions::default(), + } + }; + + match Client::new(&provider_str, opts) { + Ok(client) => { + error::clear_last_error(); + let handle = AibridgeClient::new(client); + Box::into_raw(Box::new(handle)) + } + Err(e) => { + error::set_last_error(&e); + ptr::null_mut() + } + } +} + +/// 启动客户端(初始化适配器) +/// +/// 返回 0 成功;负数为错误码(详见 `aibridge_last_error`)。 +/// +/// # Safety +/// `client` 必须是 [`aibridge_client_new`] 返回的有效指针。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_start( + client: *mut aibridge_client_t, +) -> aibridge_status_t { + catch_unwind_status(AssertUnwindSafe(|| { + if client.is_null() { + return error::set_ffi_error_simple("client 句柄为空"); + } + // SAFETY: 调用方保证 client 来自 aibridge_client_new + let handle: &AibridgeClient = &*client; + let inner = handle.arc(); + let result = block_on(async { + let mut guard = inner.lock().await; + guard.start().await + }); + match result { + Ok(()) => { + error::clear_last_error(); + AIBRIDGE_OK + } + Err(e) => error::set_last_error(&e), + } + })) +} + +/// 释放客户端句柄 +/// +/// # Safety +/// `client` 必须是 [`aibridge_client_new`] 返回的指针,且只能释放一次。 +/// 传 `nullptr` 是安全的 no-op。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_destroy(client: *mut aibridge_client_t) { + let _ = catch_unwind(AssertUnwindSafe(|| { + handle::destroy_client_impl(client); + })); +} + +// ========================================================================= +// 阻塞式调用:chat / speech +// ========================================================================= + +/// 文本对话(阻塞) +/// +/// `request_json` 为 `ChatRequest` 的 JSON 序列化字符串。 +/// 成功时 `*out_response_json` 写入 `ChatCompletion` 的 JSON,调用方需通过 +/// [`aibridge_string_free`] 释放。 +/// +/// # Safety +/// - `client` 必须来自 [`aibridge_client_new`] +/// - `request_json` 必须是合法 JSON C 字符串 +/// - `out_response_json` 必须指向可写的 `*mut c_char` 槽位 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_chat( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_response_json: *mut *mut c_char, +) -> aibridge_status_t { + catch_unwind_status(AssertUnwindSafe(|| { + chat_impl(client, request_json, out_response_json) + })) +} + +/// `aibridge_client_chat` 的实现 +/// +/// # Safety +/// `client` 必须来自 [`aibridge_client_new`](可为 null,内部校验)。 +unsafe fn chat_impl( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_response_json: *mut *mut c_char, +) -> aibridge_status_t { + if client.is_null() { + return error::set_ffi_error_simple("client 句柄为空"); + } + if out_response_json.is_null() { + return error::set_ffi_error_simple("out_response_json 输出指针为空"); + } + let req: ChatRequest = match parse_request_json(request_json, "ChatRequest") { + Ok(r) => r, + Err(status) => return status, + }; + + // SAFETY: 调用方保证 client 来自 aibridge_client_new + let handle: &AibridgeClient = &*client; + let inner = handle.arc(); + let result = block_on(async { + let guard = inner.lock().await; + guard.chat(req).await + }); + + match result { + Ok(completion) => match serde_json::to_string(&completion) { + Ok(json) => { + let ptr = alloc_cstring(json); + if ptr.is_null() { + return error::set_ffi_error_simple("响应 JSON 分配失败"); + } + *out_response_json = ptr; + error::clear_last_error(); + AIBRIDGE_OK + } + Err(e) => error::set_ffi_error( + "序列化 ChatCompletion 失败", + serde_json::json!({ "error": e.to_string() }), + ), + }, + Err(e) => error::set_last_error(&e), + } +} + +/// 文字转语音(阻塞,二进制载荷走 `aibridge_bytes_t`) +/// +/// `request_json` 为 `SpeechRequest` 的 JSON 序列化字符串。 +/// 成功时 `*out_audio` 写入二进制音频缓冲(`aibridge_bytes_t`), +/// `*out_meta_json` 写入 `SpeechResult`(不含 audio_data)的 JSON。 +/// 两者均由调用方分别通过 [`aibridge_bytes_free`] / [`aibridge_string_free`] 释放。 +/// +/// 若 Provider 仅返回 `audio_base64`/`audio_url` 而无二进制数据, +/// `*out_audio` 将为 `nullptr`(meta_json 仍写入)。 +/// +/// # Safety +/// - `client` / `request_json` / `out_audio` / `out_meta_json` 均需合法非空 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_speech( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_audio: *mut *mut aibridge_bytes_t, + out_meta_json: *mut *mut c_char, +) -> aibridge_status_t { + catch_unwind_status(AssertUnwindSafe(|| { + speech_impl(client, request_json, out_audio, out_meta_json) + })) +} + +/// `aibridge_client_speech` 的实现 +/// +/// # Safety +/// `client` 必须来自 [`aibridge_client_new`](可为 null,内部校验)。 +unsafe fn speech_impl( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_audio: *mut *mut aibridge_bytes_t, + out_meta_json: *mut *mut c_char, +) -> aibridge_status_t { + if client.is_null() { + return error::set_ffi_error_simple("client 句柄为空"); + } + if out_audio.is_null() { + return error::set_ffi_error_simple("out_audio 输出指针为空"); + } + if out_meta_json.is_null() { + return error::set_ffi_error_simple("out_meta_json 输出指针为空"); + } + let req: SpeechRequest = match parse_request_json(request_json, "SpeechRequest") { + Ok(r) => r, + Err(status) => return status, + }; + + // SAFETY: 调用方保证 client 来自 aibridge_client_new + let handle: &AibridgeClient = &*client; + let inner = handle.arc(); + let result = block_on(async { + let guard = inner.lock().await; + guard.speech(req).await + }); + + match result { + Ok(speech) => { + // 写入二进制音频(若有) + let audio_bytes = speech.get_audio_bytes(); + let audio_ptr = match audio_bytes { + Some(data) if !data.is_empty() => { + Box::into_raw(Box::new(aibridge_bytes_t::from_vec(data))) + } + _ => ptr::null_mut(), + }; + *out_audio = audio_ptr; + + // 写入 meta JSON(SpeechResult,audio_data 被 serde skip) + match serde_json::to_string(&speech) { + Ok(json) => { + let ptr = alloc_cstring(json); + if ptr.is_null() { + // audio 已分配,需回滚释放避免泄漏 + if !audio_ptr.is_null() { + handle::destroy_bytes_ptr(audio_ptr); + } + return error::set_ffi_error_simple("meta JSON 分配失败"); + } + *out_meta_json = ptr; + error::clear_last_error(); + AIBRIDGE_OK + } + Err(e) => { + if !audio_ptr.is_null() { + handle::destroy_bytes_ptr(audio_ptr); + } + error::set_ffi_error( + "序列化 SpeechResult 失败", + serde_json::json!({ "error": e.to_string() }), + ) + } + } + } + Err(e) => error::set_last_error(&e), + } +} + +// ========================================================================= +// 流式:chat_stream / stream_next / stream_destroy +// ========================================================================= + +/// 流式文本对话(创建 stream 句柄) +/// +/// `request_json` 为 `ChatRequest` 的 JSON 序列化字符串。 +/// 成功时 `*out_stream` 写入 stream 句柄,调用方通过 +/// [`aibridge_stream_next`] 拉取 chunk,最后用 [`aibridge_stream_destroy`] 释放。 +/// +/// # Safety +/// - `client` / `request_json` / `out_stream` 均需合法非空 +#[no_mangle] +pub unsafe extern "C" fn aibridge_client_chat_stream( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_stream: *mut *mut aibridge_stream_t, +) -> aibridge_status_t { + catch_unwind_status(AssertUnwindSafe(|| { + chat_stream_impl(client, request_json, out_stream) + })) +} + +/// `aibridge_client_chat_stream` 的实现 +/// +/// # Safety +/// `client` 必须来自 [`aibridge_client_new`](可为 null,内部校验)。 +unsafe fn chat_stream_impl( + client: *mut aibridge_client_t, + request_json: *const c_char, + out_stream: *mut *mut aibridge_stream_t, +) -> aibridge_status_t { + if client.is_null() { + return error::set_ffi_error_simple("client 句柄为空"); + } + if out_stream.is_null() { + return error::set_ffi_error_simple("out_stream 输出指针为空"); + } + let req: ChatRequest = match parse_request_json(request_json, "ChatRequest") { + Ok(r) => r, + Err(status) => return status, + }; + + // SAFETY: 调用方保证 client 来自 aibridge_client_new + let handle: &AibridgeClient = &*client; + let inner = handle.arc(); + let result = block_on(async { + let guard = inner.lock().await; + guard.chat_stream(req).await + }); + + match result { + Ok(stream) => { + let stream_handle = crate::handle::AibridgeStream::new(stream); + *out_stream = Box::into_raw(Box::new(stream_handle)); + error::clear_last_error(); + AIBRIDGE_OK + } + Err(e) => error::set_last_error(&e), + } +} + +/// 拉取下一个流式 chunk(阻塞) +/// +/// 返回: +/// - `0`(`AIBRIDGE_STREAM_CHUNK`):拉到一个 chunk,`*out_chunk_json` 写入 JSON +/// - `1`(`AIBRIDGE_STREAM_END`):流正常结束 +/// - 负数:错误(`aibridge_last_error` 已写入) +/// +/// # Safety +/// - `stream` 必须来自 [`aibridge_client_chat_stream`] +/// - `out_chunk_json` 必须指向可写的 `*mut c_char` 槽位 +#[no_mangle] +pub unsafe extern "C" fn aibridge_stream_next( + stream: *mut aibridge_stream_t, + out_chunk_json: *mut *mut c_char, +) -> aibridge_status_t { + catch_unwind_status(AssertUnwindSafe(|| { + stream::stream_next_impl(stream, out_chunk_json) + })) +} + +/// 释放 stream 句柄(触发 Rust drop → tokio task abort) +/// +/// # Safety +/// `stream` 必须来自 [`aibridge_client_chat_stream`],且只能释放一次。 +/// 传 `nullptr` 是安全的 no-op。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_stream_destroy(stream: *mut aibridge_stream_t) { + let _ = catch_unwind(AssertUnwindSafe(|| { + handle::destroy_stream_impl(stream); + })); +} + +// ========================================================================= +// 错误查询 +// ========================================================================= + +/// 读取当前线程的 last_error(JSON 字符串) +/// +/// 返回指向线程局部 `CString` 内部缓冲的指针,**调用方不应释放**。 +/// 若当前线程无错误,返回 `nullptr`。 +/// +/// 输出格式:`{"code":"...","message":"...","details":...,"retryable":bool}` +/// +/// # Safety +/// 本函数本身无内存不安全操作(仅读取线程局部变量并返回裸指针), +/// 标 `unsafe extern "C"` 仅为符合 FFI 调用约定。调用方**不得**释放返回的指针, +/// 且需注意返回的指针仅在当前线程的下一次 FFI 调用前保证有效(thread_local 语义)。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_last_error() -> *const c_char { + // 该函数本身不 panic(仅读 thread_local + 返回指针),但仍兜底 + let result = catch_unwind(AssertUnwindSafe(error::last_error_ptr)); + match result { + Ok(ptr) => ptr, + Err(_) => ptr::null(), + } +} + +// ========================================================================= +// 内部辅助 +// ========================================================================= + +/// 解析请求 JSON 字符串为目标类型 +/// +/// 返回 `Ok(parsed)` 或 `Err(ffi_status)`(last_error 已写入)。 +fn parse_request_json( + request_json: *const c_char, + type_name: &str, +) -> Result { + let json_str = match cstr_to_string(request_json) { + Some(Ok(s)) => s, + Some(Err(_)) => { + return Err(error::set_ffi_error( + &format!("{type_name} 请求 JSON 不是合法 UTF-8"), + serde_json::Value::Null, + )); + } + None => { + return Err(error::set_ffi_error( + &format!("{type_name} 请求 JSON 为空指针"), + serde_json::Value::Null, + )); + } + }; + match serde_json::from_str::(&json_str) { + Ok(v) => Ok(v), + Err(e) => Err(error::set_ffi_error( + &format!("{type_name} 请求 JSON 反序列化失败"), + serde_json::json!({ "error": e.to_string() }), + )), + } +} + +/// 包裹 catch_unwind:把 panic 转为 `AIBRIDGE_ERR_FFI` 并写入 last_error +fn catch_unwind_status(f: F) -> aibridge_status_t +where + F: FnOnce() -> aibridge_status_t, +{ + match catch_unwind(AssertUnwindSafe(f)) { + Ok(status) => status, + Err(_) => error::set_ffi_error_simple("FFI 调用内部 panic"), + } +} + +// ========================================================================= +// 单元测试 +// ========================================================================= +#[cfg(test)] +mod tests { + use super::*; + use crate::error::AIBRIDGE_ERR_FFI; + use std::ffi::CString; + + /// 构造 C 字符串指针的辅助 + fn cstr(s: &str) -> *const c_char { + CString::new(s).unwrap().into_raw() + } + + /// 释放 `cstr` 分配的指针 + unsafe fn free_cstr(p: *const c_char) { + if !p.is_null() { + drop(CString::from_raw(p as *mut c_char)); + } + } + + #[test] + fn client_new_with_null_provider_returns_null() { + unsafe { + let ptr = aibridge_client_new(std::ptr::null(), std::ptr::null()); + assert!(ptr.is_null()); + assert!(!aibridge_last_error().is_null()); + error::clear_last_error(); + } + } + + #[test] + fn client_new_with_invalid_config_json_returns_null() { + unsafe { + let provider = cstr("openai"); + let bad = cstr("{not json}"); + let ptr = aibridge_client_new(provider, bad); + assert!(ptr.is_null()); + let err_ptr = aibridge_last_error(); + assert!(!err_ptr.is_null()); + let json = crate::string::cstr_to_string_unchecked(err_ptr); + assert!(json.contains("config_json")); + error::clear_last_error(); + free_cstr(provider); + free_cstr(bad); + } + } + + #[test] + fn client_new_with_unknown_provider_returns_null_and_provider_not_found() { + // 阶段 0.4:工厂占位,任何 provider 都返 ProviderNotFound + unsafe { + let provider = cstr("nonexistent_provider"); + let config = cstr(r#"{"api_key":"sk-test","timeout":60}"#); + let ptr = aibridge_client_new(provider, config); + assert!(ptr.is_null()); + let err_ptr = aibridge_last_error(); + assert!(!err_ptr.is_null()); + let json = crate::string::cstr_to_string_unchecked(err_ptr); + let v: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(v["code"], "provider_not_found"); + assert_eq!(v["retryable"], false); + error::clear_last_error(); + free_cstr(provider); + free_cstr(config); + } + } + + #[test] + fn client_new_missing_api_key_returns_validation_error() { + // openai 缺 api_key:core 会返 ValidationError(validate 在 create_adapter 前) + unsafe { + let provider = cstr("openai"); + let ptr = aibridge_client_new(provider, std::ptr::null()); + assert!(ptr.is_null()); + let err_ptr = aibridge_last_error(); + assert!(!err_ptr.is_null()); + let json = crate::string::cstr_to_string_unchecked(err_ptr); + let v: serde_json::Value = serde_json::from_str(&json).unwrap(); + assert_eq!(v["code"], "validation_error"); + error::clear_last_error(); + free_cstr(provider); + } + } + + #[test] + fn chat_with_null_client_returns_ffi_error() { + unsafe { + let req = cstr(r#"{"model":"gpt-4o","messages":[]}"#); + let mut out: *mut c_char = std::ptr::null_mut(); + let status = aibridge_client_chat(std::ptr::null_mut(), req, &mut out); + assert_eq!(status, AIBRIDGE_ERR_FFI); + assert!(out.is_null()); + error::clear_last_error(); + free_cstr(req); + } + } + + #[test] + fn chat_with_null_out_returns_ffi_error() { + unsafe { + let req = cstr(r#"{"model":"gpt-4o","messages":[]}"#); + // 构造非空但无效的 client 指针(仅用于触发 out 校验,不实际解引用) + let fake_client = 0x1usize as *mut aibridge_client_t; + let status = aibridge_client_chat(fake_client, req, std::ptr::null_mut()); + assert_eq!(status, AIBRIDGE_ERR_FFI); + error::clear_last_error(); + free_cstr(req); + } + } + + #[test] + fn chat_with_invalid_request_json_returns_ffi_error() { + unsafe { + // 用非空 client 指针,但 request_json 非法 + let fake_client = 0x1usize as *mut aibridge_client_t; + let bad = cstr("{bad json}"); + let mut out: *mut c_char = std::ptr::null_mut(); + let status = aibridge_client_chat(fake_client, bad, &mut out); + assert_eq!(status, AIBRIDGE_ERR_FFI); + let err_ptr = aibridge_last_error(); + let json = crate::string::cstr_to_string_unchecked(err_ptr); + assert!(json.contains("ChatRequest")); + error::clear_last_error(); + free_cstr(bad); + } + } + + #[test] + fn speech_with_null_client_returns_ffi_error() { + unsafe { + let req = cstr(r#"{"model":"tts-1","input":"hi","voice":"alloy"}"#); + let mut out_audio: *mut aibridge_bytes_t = std::ptr::null_mut(); + let mut out_meta: *mut c_char = std::ptr::null_mut(); + let status = + aibridge_client_speech(std::ptr::null_mut(), req, &mut out_audio, &mut out_meta); + assert_eq!(status, AIBRIDGE_ERR_FFI); + assert!(out_audio.is_null()); + assert!(out_meta.is_null()); + error::clear_last_error(); + free_cstr(req); + } + } + + #[test] + fn speech_with_null_out_pointers_returns_ffi_error() { + unsafe { + let req = cstr(r#"{"model":"tts-1","input":"hi","voice":"alloy"}"#); + let fake_client = 0x1usize as *mut aibridge_client_t; + // out_audio 为空 + let mut out_meta: *mut c_char = std::ptr::null_mut(); + let status = + aibridge_client_speech(fake_client, req, std::ptr::null_mut(), &mut out_meta); + assert_eq!(status, AIBRIDGE_ERR_FFI); + error::clear_last_error(); + free_cstr(req); + } + } + + #[test] + fn chat_stream_with_null_client_returns_ffi_error() { + unsafe { + let req = cstr(r#"{"model":"gpt-4o","messages":[]}"#); + let mut out_stream: *mut aibridge_stream_t = std::ptr::null_mut(); + let status = aibridge_client_chat_stream(std::ptr::null_mut(), req, &mut out_stream); + assert_eq!(status, AIBRIDGE_ERR_FFI); + assert!(out_stream.is_null()); + error::clear_last_error(); + free_cstr(req); + } + } + + #[test] + fn stream_next_null_stream_returns_ffi_error() { + unsafe { + let mut out: *mut c_char = std::ptr::null_mut(); + let status = aibridge_stream_next(std::ptr::null_mut(), &mut out); + assert_eq!(status, AIBRIDGE_ERR_FFI); + error::clear_last_error(); + } + } + + #[test] + fn stream_destroy_null_is_noop() { + unsafe { + aibridge_stream_destroy(std::ptr::null_mut()); + } + } + + #[test] + fn client_destroy_null_is_noop() { + unsafe { + aibridge_client_destroy(std::ptr::null_mut()); + } + } + + #[test] + fn last_error_returns_null_after_clear() { + error::clear_last_error(); + unsafe { + assert!(aibridge_last_error().is_null()); + } + } + + #[test] + fn parse_request_json_handles_null_pointer() { + let status = parse_request_json::(std::ptr::null(), "ChatRequest"); + assert!(status.is_err()); + assert_eq!(status.unwrap_err(), AIBRIDGE_ERR_FFI); + error::clear_last_error(); + } + + #[test] + fn parse_request_json_handles_valid_json() { + let json = r#"{"model":"gpt-4o","messages":[]}"#; + let c = CString::new(json).unwrap(); + let result = parse_request_json::(c.as_ptr(), "ChatRequest"); + assert!(result.is_ok()); + let req = result.unwrap(); + assert_eq!(req.model, "gpt-4o"); + } + + #[test] + fn alloc_cstring_roundtrip_via_string_free() { + let ptr = alloc_cstring("test".to_string()); + assert!(!ptr.is_null()); + unsafe { + aibridge_string_free(ptr); + } + } + + #[test] + fn catch_unwind_status_converts_panic_to_ffi_error() { + let status = catch_unwind_status(|| panic!("boom")); + assert_eq!(status, AIBRIDGE_ERR_FFI); + error::clear_last_error(); + } +} diff --git a/crates/aibridge-ffi/src/runtime.rs b/crates/aibridge-ffi/src/runtime.rs new file mode 100644 index 0000000..c805d03 --- /dev/null +++ b/crates/aibridge-ffi/src/runtime.rs @@ -0,0 +1,51 @@ +//! 全局 tokio 运行时 +//! +//! 为所有 FFI 调用提供共享的多线程异步运行时。 +//! C 侧可并发调用不同 client;每个 FFI 函数内部通过 [`block_on`] 在该 runtime 上 +//! 阻塞执行异步 future。 +//! +//! 设计要点(与设计文档 7 节一致): +//! - `once_cell::Lazy` 全局单例,多线程(`multi_thread` + `worker_threads`) +//! - [`block_on`] 辅助函数:`runtime.handle().block_on(future)` +//! - 不暴露 handle 给 C 侧,所有调度由 Rust 内部完成 + +use once_cell::sync::Lazy; +use tokio::runtime::Runtime; + +/// 全局 tokio 运行时单例 +/// +/// 多线程运行时,工作线程数取 tokio 默认值(= CPU 核心数)。 +/// 使用 `Lazy` 保证首次访问时初始化,且线程安全。 +static RUNTIME: Lazy = Lazy::new(|| Runtime::new().expect("初始化 tokio 运行时失败")); + +/// 在全局 runtime 上阻塞执行一个 future 并返回其结果 +/// +/// 所有 FFI 函数内部异步逻辑都通过此函数驱动。 +/// 使用 `Runtime::block_on`(而非 `Handle::block_on`),因此**不要求** future 为 `Send`, +/// 这样允许在 future 内持有非 `Send` 的 `MutexGuard` 跨 `.await`。 +pub fn block_on(future: F) -> F::Output +where + F: std::future::Future, +{ + RUNTIME.block_on(future) +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn block_on_executes_future() { + let result = block_on(async { 42 }); + assert_eq!(result, 42); + } + + #[test] + fn runtime_is_shared_across_calls() { + // 多次调用应复用同一 runtime(Lazy 只初始化一次) + let a = block_on(async { "a" }); + let b = block_on(async { "b" }); + assert_eq!(a, "a"); + assert_eq!(b, "b"); + } +} diff --git a/crates/aibridge-ffi/src/stream.rs b/crates/aibridge-ffi/src/stream.rs new file mode 100644 index 0000000..8ffbd8a --- /dev/null +++ b/crates/aibridge-ffi/src/stream.rs @@ -0,0 +1,200 @@ +//! 流式 chunk 拉取 +//! +//! 对应设计文档 7 节的流式接口: +//! - [`stream_next_impl`]:阻塞拉取下一个 chunk +//! - 返回 `AIBRIDGE_STREAM_CHUNK`(0) = 拉到一个 chunk(写入 out_chunk_json) +//! - 返回 `AIBRIDGE_STREAM_END`(1) = 流正常结束 +//! - 返回负数 = 流错误(last_error 已写入) +//! +//! 实现要点:在全局 runtime 上 `block_on` 拉取 stream 的 next(), +//! 将 `ChatCompletionChunk` 序列化为 JSON 写入 out 指针。 + +use crate::error::{self, AIBRIDGE_STREAM_CHUNK, AIBRIDGE_STREAM_END}; +use crate::handle::{aibridge_stream_t, AibridgeStream}; +use crate::string::alloc_cstring; +use futures::stream::StreamExt; +use std::os::raw::c_char; + +/// 阻塞拉取下一个流式 chunk +/// +/// # 返回 +/// - `AIBRIDGE_STREAM_CHUNK`(0):成功拉取一个 chunk,`*out_chunk_json` 写入 JSON +/// - `AIBRIDGE_STREAM_END`(1):流正常结束 +/// - 负数:流错误(`aibridge_last_error` 已写入线程局部槽) +/// +/// # Safety +/// - `stream` 必须是 [`crate::aibridge_client_chat_stream`] 返回的有效指针 +/// - `out_chunk_json` 必须指向可写的 `*mut c_char` 槽位 +/// - 输出的字符串由调用方通过 [`crate::aibridge_string_free`] 释放 +pub(crate) unsafe fn stream_next_impl( + stream: *mut aibridge_stream_t, + out_chunk_json: *mut *mut c_char, +) -> i32 { + // —— 参数校验 —— + if stream.is_null() { + return error::set_ffi_error_simple("stream 句柄为空"); + } + if out_chunk_json.is_null() { + return error::set_ffi_error_simple("out_chunk_json 输出指针为空"); + } + + // SAFETY: 调用方保证 stream 来自 chat_stream,且未被释放 + let stream_handle: &mut AibridgeStream = &mut *stream; + + // 已结束的流直接返回 EOF + if stream_handle.ended { + return AIBRIDGE_STREAM_END; + } + + // 取出 stream 的 next(需要 &mut,stream 字段为 Option) + let mut stream_opt = stream_handle.stream.take(); + let next_result = match stream_opt.as_mut() { + Some(s) => crate::runtime::block_on(s.next()), + None => { + // stream 已被取走(不应发生,但防御性处理) + stream_handle.stream = stream_opt; + stream_handle.ended = true; + return AIBRIDGE_STREAM_END; + } + }; + + match next_result { + None => { + // 流正常结束 + stream_handle.ended = true; + stream_handle.stream = stream_opt; + AIBRIDGE_STREAM_END + } + Some(Ok(chunk)) => { + // 序列化 chunk 为 JSON + match serde_json::to_string(&chunk) { + Ok(json) => { + let ptr = alloc_cstring(json); + if ptr.is_null() { + stream_handle.stream = stream_opt; + return error::set_ffi_error_simple("chunk JSON 分配失败"); + } + *out_chunk_json = ptr; + stream_handle.stream = stream_opt; + AIBRIDGE_STREAM_CHUNK + } + Err(e) => { + stream_handle.stream = stream_opt; + error::set_ffi_error( + "序列化 ChatCompletionChunk 失败", + serde_json::json!({ "error": e.to_string() }), + ) + } + } + } + Some(Err(err)) => { + // 流中产生错误:写入 last_error 并返回对应错误码 + stream_handle.ended = true; + stream_handle.stream = stream_opt; + error::set_last_error(&err) + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::error::AIBRIDGE_ERR_FFI; + use crate::handle::AibridgeStream; + use aibridge_core::error::AibridgeError; + use aibridge_core::model::ChatCompletionChunk; + use futures::stream; + use std::ffi::CString; + + /// 构造一个测试 stream 句柄:从 chunk JSON 数组生成 + fn make_test_stream(chunks: Vec<&str>) -> *mut aibridge_stream_t { + let parsed: Vec> = chunks + .into_iter() + .map(|s| { + serde_json::from_str::(s) + .map_err(|e| AibridgeError::validation(format!("测试 fixture 解析失败: {e}"))) + }) + .collect(); + let s = stream::iter(parsed).boxed(); + let handle = AibridgeStream::new(s); + Box::into_raw(Box::new(handle)) + } + + fn free_stream(ptr: *mut aibridge_stream_t) { + unsafe { + crate::handle::destroy_stream_impl(ptr); + } + } + + #[test] + fn next_returns_chunks_then_end() { + // 用一个合法的 chunk JSON + let chunk_json = r#"{"id":"chatcmpl-1","object":"chat.completion.chunk","created":0,"model":"m","choices":[]}"#; + let ptr = make_test_stream(vec![chunk_json, chunk_json]); + unsafe { + let mut out: *mut c_char = std::ptr::null_mut(); + let s1 = stream_next_impl(ptr, &mut out); + assert_eq!(s1, AIBRIDGE_STREAM_CHUNK); + assert!(!out.is_null()); + let json = cstr_to_string_unchecked_internal(out); + assert!(json.contains("chatcmpl-1")); + crate::aibridge_string_free(out); + + let mut out2: *mut c_char = std::ptr::null_mut(); + let s2 = stream_next_impl(ptr, &mut out2); + assert_eq!(s2, AIBRIDGE_STREAM_CHUNK); + crate::aibridge_string_free(out2); + + // 第三次应返回 END + let mut out3: *mut c_char = std::ptr::null_mut(); + let s3 = stream_next_impl(ptr, &mut out3); + assert_eq!(s3, AIBRIDGE_STREAM_END); + assert!(out3.is_null()); + } + free_stream(ptr); + } + + #[test] + fn next_with_null_stream_returns_ffi_error() { + unsafe { + let mut out: *mut c_char = std::ptr::null_mut(); + let s = stream_next_impl(std::ptr::null_mut(), &mut out); + assert_eq!(s, AIBRIDGE_ERR_FFI); + // last_error 应已写入 + assert!(!error::last_error_ptr().is_null()); + error::clear_last_error(); + } + } + + #[test] + fn next_with_null_out_returns_ffi_error() { + let ptr = make_test_stream(vec![]); + unsafe { + let s = stream_next_impl(ptr, std::ptr::null_mut()); + assert_eq!(s, AIBRIDGE_ERR_FFI); + error::clear_last_error(); + } + free_stream(ptr); + } + + #[test] + fn ended_stream_returns_eof_repeatedly() { + let ptr = make_test_stream(vec![]); + unsafe { + let mut out: *mut c_char = std::ptr::null_mut(); + assert_eq!(stream_next_impl(ptr, &mut out), AIBRIDGE_STREAM_END); + // 再次拉取仍返回 END + assert_eq!(stream_next_impl(ptr, &mut out), AIBRIDGE_STREAM_END); + } + free_stream(ptr); + } + + // 内部辅助:复用 string 模块的 unchecked 转换 + fn cstr_to_string_unchecked_internal(ptr: *mut c_char) -> String { + unsafe { crate::string::cstr_to_string_unchecked(ptr) } + } + + // 抑制未使用警告(CString 在某些测试路径用到) + #[allow(dead_code)] + fn _suppress(_c: CString) {} +} diff --git a/crates/aibridge-ffi/src/string.rs b/crates/aibridge-ffi/src/string.rs new file mode 100644 index 0000000..aa1c659 --- /dev/null +++ b/crates/aibridge-ffi/src/string.rs @@ -0,0 +1,132 @@ +//! C 字符串与二进制辅助函数 +//! +//! 提供跨 FFI 边界的字符串/字节传递与释放: +//! - [`aibridge_string_free`]:释放 Rust 分配的 `char*` +//! - [`aibridge_bytes_free`]:释放 Rust 分配的 [`aibridge_bytes_t`] +//! - [`cstr_to_string`] / [`cstr_to_string_unchecked`]:从 C 指针构造 Rust String +//! - [`alloc_cstring`]:把 Rust String 装进 CString 并返回裸指针(调用方负责释放) +//! +//! 设计要点(设计文档 7 节):Rust 分配的 `char*` / `aibridge_bytes_t*` +//! 必须由调用方通过对应的 `_free` 函数释放。各语言绑定用 RAII 封装。 + +use crate::handle::aibridge_bytes_t; +use std::ffi::{CStr, CString}; +use std::os::raw::c_char; +use std::ptr; + +/// 把 C `*const c_char` 转为 Rust `String` +/// +/// 空指针返回 `None`;非 UTF-8 返回 `Err`。 +pub fn cstr_to_string(ptr: *const c_char) -> Option> { + if ptr.is_null() { + return None; + } + // SAFETY: 调用方保证 ptr 指向合法的 NUL 结尾 C 字符串 + let cstr = unsafe { CStr::from_ptr(ptr) }; + Some(cstr.to_str().map(|s| s.to_string())) +} + +/// 把 C `*const c_char` 转为 Rust `String`(不检查 UTF-8,仅供内部测试用) +/// +/// # Safety +/// `ptr` 必须指向合法的 NUL 结尾 C 字符串且为有效 UTF-8。 +#[allow(dead_code)] +pub unsafe fn cstr_to_string_unchecked(ptr: *const c_char) -> String { + if ptr.is_null() { + return String::new(); + } + CStr::from_ptr(ptr).to_string_lossy().into_owned() +} + +/// 把 Rust `String` 装进 `CString` 并返回裸指针 +/// +/// 调用方需通过 [`aibridge_string_free`] 释放。返回 `nullptr` 表示分配失败 +/// 或字符串含内嵌 NUL。 +pub fn alloc_cstring(s: String) -> *mut c_char { + match CString::new(s) { + Ok(cstr) => cstr.into_raw(), + Err(_) => ptr::null_mut(), + } +} + +/// 释放 Rust 分配的 C 字符串 +/// +/// # Safety +/// `ptr` 必须是由 [`alloc_cstring`](或 `CString::into_raw`)分配的指针, +/// 且只能释放一次。传 `nullptr` 是安全的 no-op。 +/// +/// 对应设计文档的 `aibridge_string_free`。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_string_free(ptr: *mut c_char) { + if !ptr.is_null() { + // SAFETY: 调用方保证 ptr 来自 CString::into_raw,且未释放过 + drop(CString::from_raw(ptr)); + } +} + +/// 释放 Rust 分配的二进制缓冲 +/// +/// # Safety +/// `ptr` 必须是指向由 FFI 函数(如 `aibridge_client_speech`)输出的 +/// `aibridge_bytes_t*` 指针,且只能释放一次。传 `nullptr` 是安全的 no-op。 +/// +/// 对应设计文档的 `aibridge_bytes_free`。 +#[no_mangle] +pub unsafe extern "C" fn aibridge_bytes_free(ptr: *mut aibridge_bytes_t) { + if !ptr.is_null() { + // SAFETY: 调用方保证 ptr 来自 Box::into_raw,且未释放过 + let boxed = Box::from_raw(ptr); + // 还原并释放内部 [u8] 缓冲(from_vec 时用 Box<[u8]> 装箱) + if !boxed.ptr.is_null() && boxed.len > 0 { + let slice = std::slice::from_raw_parts_mut(boxed.ptr as *mut u8, boxed.len); + drop(Box::from_raw(slice)); + } + // Box 随作用域结束 drop + drop(boxed); + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn alloc_and_free_cstring_roundtrip() { + let ptr = alloc_cstring("你好, aibridge".to_string()); + assert!(!ptr.is_null()); + let s = unsafe { cstr_to_string_unchecked(ptr) }; + assert_eq!(s, "你好, aibridge"); + unsafe { aibridge_string_free(ptr) }; + } + + #[test] + fn alloc_cstring_with_embedded_nul_returns_null() { + // 含内嵌 NUL 的字符串无法构造 CString + let ptr = alloc_cstring("a\0b".to_string()); + assert!(ptr.is_null()); + } + + #[test] + fn string_free_null_is_noop() { + unsafe { aibridge_string_free(std::ptr::null_mut()) }; + } + + #[test] + fn cstr_to_string_null_returns_none() { + assert!(cstr_to_string(std::ptr::null()).is_none()); + } + + #[test] + fn bytes_free_null_is_noop() { + unsafe { aibridge_bytes_free(std::ptr::null_mut()) }; + } + + #[test] + fn bytes_alloc_and_free_roundtrip() { + // 构造一个 aibridge_bytes_t,释放应不泄漏/不崩溃 + let data = vec![1u8, 2, 3, 4, 5]; + let b = Box::new(aibridge_bytes_t::from_vec(data)); + let ptr = Box::into_raw(b); + unsafe { aibridge_bytes_free(ptr) }; + } +} From 8020c9c59da3a7c96f3fe1078888a070c83e6d1d Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 11:40:41 +0800 Subject: [PATCH 06/55] =?UTF-8?q?feat(aibridge-core):=20echo=20=E9=80=82?= =?UTF-8?q?=E9=85=8D=E5=99=A8=E7=94=A8=E4=BA=8E=E9=98=B6=E6=AE=B50.6=20?= =?UTF-8?q?=E7=AE=A1=E7=BA=BF=E9=AA=8C=E8=AF=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 21 +- crates/aibridge-core/src/adapters/echo.rs | 687 ++++++++++++++++++++ crates/aibridge-core/src/adapters/mod.rs | 3 + crates/aibridge-core/src/client.rs | 36 +- 4 files changed, 744 insertions(+), 3 deletions(-) create mode 100644 crates/aibridge-core/src/adapters/echo.rs diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index 100cfc5..2c35910 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -8,13 +8,16 @@ //! - 阶段 0.4 暂只占位分支(返 ProviderNotFound),具体适配器阶段 1 起填充 use crate::adapter::Adapter; +use crate::adapters::echo::EchoAdapter; use crate::config::ProviderConfig; use crate::error::{AibridgeError, Result}; /// 已支持的 provider 列表(用于错误信息与测试) /// /// 阶段 1 起逐步填充实际支持的 provider 名。 +/// `echo` 为阶段 0.6 管线验证用的 mock 适配器,常驻可用。 pub const KNOWN_PROVIDERS: &[&str] = &[ + "echo", "openai", "agnes", "volcengine_cv", @@ -38,10 +41,13 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ /// 根据配置创建适配器实例 /// /// 对应 Python v1 `AdapterFactory.create`。 -/// 阶段 0.4:所有分支均为占位(返 ProviderNotFound),具体适配器在阶段 1 起填充。 +/// `echo` 为阶段 0.6 管线验证用 mock 适配器,已实现; +/// 其余 provider 阶段 0.4 占位(返 ProviderNotFound),阶段 1 起填充。 pub fn create_adapter(config: ProviderConfig) -> Result> { let provider = config.provider_type.as_str(); match provider { + // Echo(Mock)适配器:阶段 0.6 管线验证用,已实现 + "echo" => Ok(Box::new(EchoAdapter::new())), // 阶段 1 MVP 适配器(阶段 1.0 起填充实际构造逻辑) "openai" | "agnes" | "volcengine_cv" | "gemini" => { // TODO(阶段 1): 引入 adapters::openai::OpenAiAdapter 等具体实现 @@ -148,6 +154,7 @@ mod tests { #[test] fn is_known_provider_recognizes_known() { + assert!(is_known_provider("echo")); assert!(is_known_provider("openai")); assert!(is_known_provider("edge-tts")); assert!(is_known_provider("assemblyai")); @@ -167,4 +174,16 @@ mod tests { panic!("应为 ProviderNotFound"); } } + + #[test] + fn create_echo_returns_adapter() { + // echo 免认证,无需 api_key + let config = ProviderConfig::from_options("echo", ClientOptions::default()); + let result = create_adapter(config); + assert!(result.is_ok(), "工厂应能创建 echo 适配器"); + let adapter = result.unwrap(); + assert_eq!(adapter.provider_type(), "echo"); + assert_eq!(adapter.provider_name(), "Echo (Mock)"); + assert!(!adapter.requires_api_key()); + } } diff --git a/crates/aibridge-core/src/adapters/echo.rs b/crates/aibridge-core/src/adapters/echo.rs new file mode 100644 index 0000000..f14e5d5 --- /dev/null +++ b/crates/aibridge-core/src/adapters/echo.rs @@ -0,0 +1,687 @@ +//! Echo(Mock)适配器 +//! +//! 用于阶段 0.6 五语言管线验证的回显适配器。 +//! 不调用任何真实网络,所有方法返回固定/回显响应,便于在没有 API Key 与 +//! 外部依赖的情况下端到端验证 chat / chat_stream / image / video / embed / +//! transcribe / speech / list_models / list_voices 等能力的数据流。 +//! +//! 设计要点: +//! - `provider_type = "echo"`,`requires_api_key = false`(免认证,便于管线验证) +//! - `capabilities` 声明全部能力,使 Router / Client 的能力检查均通过 +//! - `chat` 回显最后一条 user 消息内容并追加 " [echo]" +//! - `chat_stream` 用 `async_stream` 产生 3 个 chunk,逐步拼接出完整回显 +//! - 二进制能力(speech)返回固定字节,验证跨语言二进制传递 + +use std::time::{SystemTime, UNIX_EPOCH}; + +use async_trait::async_trait; +use futures::stream::StreamExt; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::error::Result; +use crate::model::audio::{SpeechResult, TranscriptionResult}; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChoiceMessage, DeltaMessage, UserContent, +}; +use crate::model::common::{ModelInfo, ModelType, TaskStatus, VoiceInfo}; +use crate::model::image::{ImageData, ImageRequest, ImageResult}; +use crate::model::options::{ + EmbedInput, EmbedRequest, EmbeddingItem, EmbeddingResult, EmbeddingUsage, EmbeddingVector, +}; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; + +/// Echo(Mock)适配器 +/// +/// 不持有任何资源(无 HTTP 客户端),构造轻量。 +/// 所有方法返回固定/回显响应,不调用网络。 +#[derive(Debug, Default, Clone)] +pub struct EchoAdapter; + +impl EchoAdapter { + /// 创建 Echo 适配器 + /// + /// 无需任何配置,构造不进行网络初始化。 + pub fn new() -> Self { + Self + } + + /// 构造全部能力集合(echo 声明支持所有能力,便于管线验证) + fn all_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::ToolCall); + caps.insert(Capabilities::Reasoning); + caps.insert(Capabilities::JsonMode); + caps.insert(Capabilities::WebSearch); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::ImageEdit); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::AudioTranscribe); + caps.insert(Capabilities::AudioSpeech); + caps.insert(Capabilities::ListVoices); + caps + } + + /// 提取请求中最后一条 user 消息的文本内容 + /// + /// 多模态消息取首个文本部件;无 user 消息时返回空串。 + fn last_user_text(req: &ChatRequest) -> String { + for msg in req.messages.iter().rev() { + if let ChatMessage::User { content, .. } = msg { + return match content { + UserContent::Text(s) => s.clone(), + UserContent::Parts(parts) => parts + .iter() + .find_map(|p| match p { + crate::model::chat::ContentPart::Text { text } => Some(text.clone()), + _ => None, + }) + .unwrap_or_default(), + }; + } + } + String::new() + } + + /// 当前 Unix 时间戳(秒) + fn now_secs() -> u64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs()) + .unwrap_or(0) + } + + /// 构造固定模型列表 + fn fixed_models() -> Vec { + vec![ + ModelInfo { + id: "echo-chat".into(), + name: "Echo Chat".into(), + model_type: ModelType::Chat, + provider: "echo".into(), + capabilities: vec!["chat".into(), "chat_stream".into()], + max_tokens: Some(4096), + supports_streaming: true, + description: Some("Echo mock chat model".into()), + created: None, + }, + ModelInfo { + id: "echo-image".into(), + name: "Echo Image".into(), + model_type: ModelType::Image, + provider: "echo".into(), + capabilities: vec!["text2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Echo mock image model".into()), + created: None, + }, + ModelInfo { + id: "echo-video".into(), + name: "Echo Video".into(), + model_type: ModelType::Video, + provider: "echo".into(), + capabilities: vec!["text2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Echo mock video model".into()), + created: None, + }, + ModelInfo { + id: "echo-tts".into(), + name: "Echo TTS".into(), + model_type: ModelType::Audio, + provider: "echo".into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Echo mock TTS model".into()), + created: None, + }, + ModelInfo { + id: "echo-asr".into(), + name: "Echo ASR".into(), + model_type: ModelType::Audio, + provider: "echo".into(), + capabilities: vec!["audio_transcribe".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Echo mock ASR model".into()), + created: None, + }, + ModelInfo { + id: "echo-embed".into(), + name: "Echo Embed".into(), + model_type: ModelType::Chat, + provider: "echo".into(), + capabilities: vec!["embedding".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Echo mock embedding model".into()), + created: None, + }, + ] + } + + /// 构造固定音色列表 + fn fixed_voices() -> Vec { + vec![ + VoiceInfo::builder() + .short_name("echo-voice-1") + .name("Echo Voice 1") + .locale("zh-CN") + .gender("Female") + .voice_id("echo-voice-1") + .build(), + VoiceInfo::builder() + .short_name("echo-voice-2") + .name("Echo Voice 2") + .locale("en-US") + .gender("Male") + .voice_id("echo-voice-2") + .build(), + ] + } + + /// 1x1 透明 PNG 的 Base64 编码 + /// + /// 用于 image_generate 返回的占位图,避免引入图像库。 + const PLACEHOLDER_PNG_B64: &'static str = + "iVBORw0KGgoAAAANSUhEUgAAAAEAAAABCAQAAAC1HAwCAAAAC0lEQVR42mNkYAAAAAYAAjCB0C8AAAAASUVORK5CYII="; + + /// 构造完整回显文本 + fn echo_text(req: &ChatRequest) -> String { + format!("{} [echo]", Self::last_user_text(req)) + } +} + +#[async_trait] +impl Adapter for EchoAdapter { + fn provider_type(&self) -> &str { + "echo" + } + + fn provider_name(&self) -> &str { + "Echo (Mock)" + } + + fn capabilities(&self) -> CapabilitySet { + Self::all_capabilities() + } + + /// 免认证:echo 不需要 API Key,便于管线验证 + fn requires_api_key(&self) -> bool { + false + } + + async fn start(&mut self) -> Result<()> { + // 无资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // 无资源需释放 + Ok(()) + } + + /// 文本对话:回显最后一条 user 消息内容 + " [echo]" + async fn chat(&self, req: ChatRequest) -> Result { + let content = Self::echo_text(&req); + Ok(ChatCompletion { + id: "echo-chat-1".into(), + object: "chat.completion".into(), + created: Self::now_secs(), + model: req.model, + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".into(), + content: Some(content), + tool_calls: None, + }, + finish_reason: Some("stop".into()), + }], + usage: Some(crate::model::chat::ChatUsage { + prompt_tokens: 1, + completion_tokens: 1, + total_tokens: 2, + }), + service_tier: None, + system_fingerprint: None, + }) + } + + /// 流式文本对话:用 async_stream 产生 3 个 chunk + /// + /// chunk 内容逐步拼接:第 1 块含 role,第 2 块含前半段内容, + /// 第 3 块含后半段内容并标记 finish_reason。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + let model = req.model.clone(); + let full = Self::echo_text(&req); + let created = Self::now_secs(); + + // 将完整回显拆为两段,模拟流式增量 + let mid = full.len() / 2; + let (first_half, second_half) = if full.is_empty() { + (String::new(), String::new()) + } else { + (full[..mid].to_string(), full[mid..].to_string()) + }; + + let stream = async_stream::stream! { + // chunk 1:角色声明 + yield Ok(ChatCompletionChunk { + id: "echo-stream-1".into(), + object: "chat.completion.chunk".into(), + created, + model: model.clone(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".into()), + content: None, + tool_calls: None, + }, + finish_reason: None, + }], + usage: None, + }); + // chunk 2:前半段内容 + yield Ok(ChatCompletionChunk { + id: "echo-stream-1".into(), + object: "chat.completion.chunk".into(), + created, + model: model.clone(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: None, + content: Some(first_half), + tool_calls: None, + }, + finish_reason: None, + }], + usage: None, + }); + // chunk 3:后半段内容 + 结束 + yield Ok(ChatCompletionChunk { + id: "echo-stream-1".into(), + object: "chat.completion.chunk".into(), + created, + model, + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: None, + content: Some(second_half), + tool_calls: None, + }, + finish_reason: Some("stop".into()), + }], + usage: None, + }); + }; + + Ok(stream.boxed()) + } + + /// 图像生成:返固定 1 张占位图(1x1 PNG Base64) + async fn image_generate(&self, req: ImageRequest) -> Result { + Ok(ImageResult { + id: "echo-image-1".into(), + object: "image.generation".into(), + created: Self::now_secs(), + model: req.model, + data: vec![ImageData { + url: None, + b64_json: Some(Self::PLACEHOLDER_PNG_B64.into()), + revised_prompt: Some(req.prompt), + }], + }) + } + + /// 创建视频任务:返固定已完成任务 + async fn video_create(&self, req: VideoRequest) -> Result { + Ok(VideoTask { + task_id: "echo-task-1".into(), + model: req.model, + status: TaskStatus::Success, + created_at: Self::now_secs(), + }) + } + + /// 查询视频任务状态:返固定完成状态 + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + Ok(VideoStatus { + task_id: task_id.to_string(), + status: TaskStatus::Success, + video_url: Some("https://example.com/echo.mp4".into()), + progress: Some(100), + error: None, + created_at: Some(Self::now_secs()), + updated_at: Some(Self::now_secs()), + }) + } + + /// 文本嵌入:输入 N 条文本返 N 个 3 维向量 + async fn embed(&self, req: EmbedRequest) -> Result { + let n = match &req.input { + EmbedInput::Single(_) => 1, + EmbedInput::Multiple(texts) => texts.len(), + }; + // 每个向量用 3 维占位值,索引区分以验证数量对齐 + let data: Vec = (0..n) + .map(|i| EmbeddingItem { + object: "embedding".into(), + index: i as u32, + embedding: EmbeddingVector::Float(vec![ + i as f64, + (i as f64) + 0.1, + (i as f64) + 0.2, + ]), + }) + .collect(); + Ok(EmbeddingResult { + object: "list".into(), + data, + model: req.model, + usage: Some(EmbeddingUsage { + prompt_tokens: n as u64, + total_tokens: n as u64, + }), + }) + } + + /// 语音转文字:返固定转写文本 + async fn transcribe( + &self, + _req: crate::model::audio::TranscribeRequest, + ) -> Result { + Ok(TranscriptionResult { + text: "echo transcription".into(), + language: Some("zh".into()), + duration: Some(1.0), + segments: None, + words: None, + task: "transcribe".into(), + usage: None, + model: Some("echo-asr".into()), + }) + } + + /// 文字转语音:返固定二进制音频(验证跨语言二进制传递) + async fn speech(&self, _req: crate::model::audio::SpeechRequest) -> Result { + Ok(SpeechResult { + audio_data: Some(b"mock-audio-data".to_vec()), + audio_url: None, + audio_base64: None, + content_type: "audio/mpeg".into(), + format: "mp3".into(), + duration: Some(1.0), + model: Some("echo-tts".into()), + }) + } + + /// 模型列表:返固定 6 个 echo 模型 + async fn list_models(&self, filter: Option) -> Result> { + let models = Self::fixed_models(); + match filter { + Some(t) => Ok(models.into_iter().filter(|m| m.model_type == t).collect()), + None => Ok(models), + } + } + + /// 音色列表:返固定 2 个 echo 音色 + async fn list_voices(&self, _language: Option<&str>) -> Result> { + Ok(Self::fixed_voices()) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::adapter::factory::create_adapter; + use crate::config::{ClientOptions, ProviderConfig}; + use crate::model::audio::{SpeechRequest, TranscribeRequest}; + use crate::model::image::FileInput; + + /// 构造测试用 ProviderConfig(echo 免认证,但 config 仍可携带任意字段) + fn echo_config() -> ProviderConfig { + ProviderConfig::from_options("echo", ClientOptions::default()) + } + + #[tokio::test] + async fn chat_echoes_last_user_message() { + let adapter = EchoAdapter::new(); + let req = ChatRequest::builder( + "echo-chat", + vec![ + ChatMessage::system("you are helpful"), + ChatMessage::user("hello world"), + ], + ) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.id, "echo-chat-1"); + assert_eq!(resp.model, "echo-chat"); + assert_eq!(resp.choices.len(), 1); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("hello world [echo]") + ); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_echoes_multimodal_user_text() { + let adapter = EchoAdapter::new(); + let req = ChatRequest::builder( + "echo-chat", + vec![ChatMessage::user_multimodal(vec![ + crate::model::chat::ContentPart::Text { + text: "describe this".into(), + }, + crate::model::chat::ContentPart::ImageUrl { + image_url: crate::model::chat::ImageUrl::new("https://example.com/x.png"), + }, + ])], + ) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("describe this [echo]") + ); + } + + #[tokio::test] + async fn chat_stream_produces_three_chunks() { + let adapter = EchoAdapter::new(); + let req = ChatRequest::builder("echo-chat", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3, "应产生 3 个 chunk"); + // 第 1 块含 role + assert_eq!( + chunks[0].choices[0].delta.role.as_deref(), + Some("assistant") + ); + // 拼接第 2、3 块内容应等于完整回显 + let mut assembled = String::new(); + assembled.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + assembled.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(assembled, "hi [echo]"); + // 第 3 块标记结束 + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn speech_returns_bytes() { + let adapter = EchoAdapter::new(); + let req = SpeechRequest::builder("echo-tts", "hello", "echo-voice-1").build(); + let resp = adapter.speech(req).await.unwrap(); + assert_eq!( + resp.audio_data.as_deref(), + Some(b"mock-audio-data".as_slice()) + ); + assert_eq!(resp.content_type, "audio/mpeg"); + assert_eq!(resp.format, "mp3"); + assert_eq!(resp.model.as_deref(), Some("echo-tts")); + } + + #[tokio::test] + async fn list_models_returns_six_echo_models() { + let adapter = EchoAdapter::new(); + let models = adapter.list_models(None).await.unwrap(); + let ids: Vec<&str> = models.iter().map(|m| m.id.as_str()).collect(); + assert_eq!( + ids, + vec![ + "echo-chat", + "echo-image", + "echo-video", + "echo-tts", + "echo-asr", + "echo-embed" + ] + ); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let adapter = EchoAdapter::new(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "echo-image"); + + let audios = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audios.len(), 2); // echo-tts + echo-asr + } + + #[tokio::test] + async fn image_generate_returns_placeholder_png() { + let adapter = EchoAdapter::new(); + let req = ImageRequest::builder("echo-image", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert!(resp.data[0].b64_json.is_some()); + // 应为合法 base64(可解码) + assert!(crate::util::decode_base64(resp.data[0].b64_json.as_ref().unwrap()).is_ok()); + } + + #[tokio::test] + async fn video_create_and_poll_roundtrip() { + let adapter = EchoAdapter::new(); + let req = VideoRequest::builder("echo-video", "a cat walking").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "echo-task-1"); + assert_eq!(task.status, TaskStatus::Success); + + let status = adapter + .video_poll("echo-task-1", "echo-video") + .await + .unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/echo.mp4") + ); + assert_eq!(status.progress, Some(100)); + } + + #[tokio::test] + async fn embed_returns_n_vectors() { + let adapter = EchoAdapter::new(); + let req = EmbedRequest { + model: "echo-embed".into(), + input: EmbedInput::Multiple(vec!["a".into(), "b".into(), "c".into()]), + dimensions: None, + encoding_format: None, + user: None, + extra: std::collections::HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 3); + for (i, item) in resp.data.iter().enumerate() { + assert_eq!(item.index, i as u32); + if let EmbeddingVector::Float(v) = &item.embedding { + assert_eq!(v.len(), 3, "每个向量应为 3 维"); + } else { + panic!("应为 Float 向量"); + } + } + } + + #[tokio::test] + async fn transcribe_returns_fixed_text() { + let adapter = EchoAdapter::new(); + let req = TranscribeRequest::builder("echo-asr", FileInput::path("/tmp/a.mp3")).build(); + let resp = adapter.transcribe(req).await.unwrap(); + assert_eq!(resp.text, "echo transcription"); + } + + #[tokio::test] + async fn list_voices_returns_two() { + let adapter = EchoAdapter::new(); + let voices = adapter.list_voices(None).await.unwrap(); + assert_eq!(voices.len(), 2); + assert_eq!(voices[0].short_name.as_deref(), Some("echo-voice-1")); + } + + #[tokio::test] + async fn recommend_voices_filters_by_gender() { + let adapter = EchoAdapter::new(); + let female = adapter + .recommend_voices(None, Some("female"), 10) + .await + .unwrap(); + assert_eq!(female.len(), 1); + assert_eq!(female[0].gender.as_deref(), Some("Female")); + } + + #[tokio::test] + async fn requires_api_key_is_false() { + let adapter = EchoAdapter::new(); + assert!(!adapter.requires_api_key()); + } + + #[tokio::test] + async fn capabilities_contains_all() { + let adapter = EchoAdapter::new(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::Embedding)); + assert!(caps.contains(&Capabilities::AudioSpeech)); + assert!(caps.contains(&Capabilities::AudioTranscribe)); + assert!(caps.contains(&Capabilities::ListVoices)); + } + + #[test] + fn factory_create_echo_succeeds() { + let result = create_adapter(echo_config()); + assert!(result.is_ok(), "工厂应能创建 echo 适配器"); + let adapter = result.unwrap(); + assert_eq!(adapter.provider_type(), "echo"); + assert_eq!(adapter.provider_name(), "Echo (Mock)"); + assert!(!adapter.requires_api_key()); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = EchoAdapter::new(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } +} diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 207295f..b45fd55 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -9,3 +9,6 @@ //! - 阶段 2c(音频):edge-tts / elevenlabs / cartesia / deepgram / assemblyai //! //! 各适配器实现 `adapter::Adapter` trait,由 `adapter::create_adapter` 工厂分发。 + +/// Echo(Mock)适配器:阶段 0.6 五语言管线验证用,不调网络返固定/回显响应 +pub mod echo; diff --git a/crates/aibridge-core/src/client.rs b/crates/aibridge-core/src/client.rs index 0e95d27..7daaeee 100644 --- a/crates/aibridge-core/src/client.rs +++ b/crates/aibridge-core/src/client.rs @@ -48,8 +48,10 @@ impl Client { // 校验:requires_api_key 时必须有 api_key // 阶段 0.4 无法预知适配器的 requires_api_key(适配器未实现), - // 暂按"有则需 key"的保守策略:未知 provider 也要求 key - config.validate(true)?; + // 暂按"有则需 key"的保守策略:未知 provider 也要求 key。 + // 例外:免认证 provider(如 echo mock 适配器)跳过 key 校验,便于管线验证 + let requires_api_key = !Self::is_free_provider(provider); + config.validate(requires_api_key)?; let adapter = create_adapter(config.clone()).map_err(|e| match e { AibridgeError::ProviderNotFound { provider: p } => { @@ -64,6 +66,14 @@ impl Client { }) } + /// 判断 provider 是否为免认证(不需要 API Key) + /// + /// 阶段 0.6:`echo` 为 mock 适配器,免认证便于五语言管线验证。 + /// 阶段 2c 起将补充 edge-tts 等真实免认证 provider。 + fn is_free_provider(provider: &str) -> bool { + matches!(provider, "echo") + } + /// 启动客户端(初始化适配器) pub async fn start(&mut self) -> Result<()> { self.adapter.start().await @@ -143,6 +153,7 @@ impl Client { #[cfg(test)] mod tests { use super::*; + use crate::model::chat::{ChatMessage, ChatRequest}; use std::env; use std::sync::Mutex; @@ -191,4 +202,25 @@ mod tests { )); env::remove_var("AIBRIDGE_ENVTEST_API_KEY"); } + + #[test] + fn new_echo_without_api_key_succeeds() { + // echo 免认证:无 api_key 也能创建客户端 + let result = Client::new("echo", ClientOptions::default()); + assert!(result.is_ok(), "echo 应免认证创建客户端"); + } + + #[tokio::test] + async fn echo_client_chat_roundtrip() { + // 端到端验证:Client → EchoAdapter → chat 回显 + let mut client = Client::new("echo", ClientOptions::default()).unwrap(); + client.start().await.unwrap(); + let req = ChatRequest::builder("echo-chat", vec![ChatMessage::user("hello")]).build(); + let resp = client.chat(req).await.unwrap(); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("hello [echo]") + ); + client.close().await.unwrap(); + } } From f36ee7c99e0d8dcd275618f965847f25c3282aa2 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 11:50:36 +0800 Subject: [PATCH 07/55] =?UTF-8?q?feat(aibridge-go):=20=E9=98=B6=E6=AE=B50.?= =?UTF-8?q?6=20CGO=20=E7=BB=91=E5=AE=9A=20+=20hello=20world?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- bindings/go/aibridge.go | 227 +++++++++++++++++++++++++++++++++++ bindings/go/error.go | 101 ++++++++++++++++ bindings/go/example/hello.go | 140 +++++++++++++++++++++ bindings/go/go.mod | 3 + bindings/go/model.go | 157 ++++++++++++++++++++++++ bindings/go/stream.go | 152 +++++++++++++++++++++++ 6 files changed, 780 insertions(+) create mode 100644 bindings/go/aibridge.go create mode 100644 bindings/go/error.go create mode 100644 bindings/go/example/hello.go create mode 100644 bindings/go/go.mod create mode 100644 bindings/go/model.go create mode 100644 bindings/go/stream.go diff --git a/bindings/go/aibridge.go b/bindings/go/aibridge.go new file mode 100644 index 0000000..c72bf86 --- /dev/null +++ b/bindings/go/aibridge.go @@ -0,0 +1,227 @@ +// Package aibridge 提供 AIBridge 的 Go 绑定(CGO 调 aibridge-ffi cdylib)。 +// +// 设计要点(详见 docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md 第 7、8 节): +// - 句柄式生命周期:Client 持 *C.aibridge_client_t,Close() 调 aibridge_client_destroy +// - 复杂 struct 走 JSON 边界:Go struct <-> JSON <-> Rust serde +// - 二进制载荷走 aibridge_bytes_t:speech 返回 []byte,必须 aibridge_bytes_free +// - 错误:FFI 返回码 + aibridge_last_error() 线程局部 JSON(同线程立即读取) +// - 流式:stream 句柄 + 阻塞 stream_next(),goroutine 循环拉取 push 到 channel +// +// FFI 遗留问题处理: +// 1. last_error 线程局部:每个 FFI 调用失败后,同一线程立即读 aibridge_last_error() 转存为 Go error +// 2. stream_next 串行:同一 stream 不可并发 next(goroutine 内串行循环保证) +// 3. Rust 分配的内存必须调对应的 free:string_free / bytes_free / stream_destroy / client_destroy +// 4. client/stream 句柄必须 destroy(RAII:Client.Close + stream 用完 destroy) +package aibridge + +/* +#cgo CFLAGS: -I${SRCDIR}/../../crates/aibridge-ffi/include +#cgo LDFLAGS: -L${SRCDIR}/../../target/debug -laibridge -lm + +#include +#include "aibridge.h" + +// 错误码常量(与 aibridge.h 宏对齐,cgo 无法直接引用宏,此处通过包装函数取值) +static int aibridge_status_ok(void) { return AIBRIDGE_OK; } +static int aibridge_status_stream_chunk(void) { return AIBRIDGE_STREAM_CHUNK; } +static int aibridge_status_stream_end(void) { return AIBRIDGE_STREAM_END; } + +// 包装:aibridge_last_error 返回 *const c_char,转为 Go 可读的字符串前需在 C 侧拷贝长度。 +// 这里直接用 Go 侧 C.GoStringPtr/GoString 读取,无需 C 包装。 +*/ +import "C" + +import ( + "encoding/json" + "runtime" + "unsafe" +) + +// statusOK 成功 +const statusOK = 0 + +// statusStreamChunk 流式:拉到一个 chunk +const statusStreamChunk = 0 + +// statusStreamEnd 流式:流已结束 +const statusStreamEnd = 1 + +// Client 是 AIBridge 客户端,持有一个 FFI client 句柄。 +// +// 生命周期:NewClient 创建 -> Start 初始化适配器 -> Chat/ChatStream/Speech 调用 -> Close 释放。 +// 必须调用 Close() 释放底层 Rust 句柄,否则内存泄漏。 +type Client struct { + ptr *C.aibridge_client_t +} + +// NewClient 创建客户端(对应 aibridge_client_new) +// +// provider 为 Provider 类型(如 "echo"、"openai")。 +// optsJSON 为 ClientOptions 的 JSON 字符串,传 nil/空串表示默认配置。 +// +// 失败返回 AibridgeError(如 provider_not_found、validation_error)。 +func NewClient(provider string, optsJSON *string) (*Client, error) { + cProvider := C.CString(provider) + defer C.free(unsafe.Pointer(cProvider)) + + var cOpts *C.char + if optsJSON != nil && *optsJSON != "" { + cOpts = C.CString(*optsJSON) + defer C.free(unsafe.Pointer(cOpts)) + } + + // aibridge_client_new 失败时返回 nullptr,错误写入 last_error(同线程) + ptr := C.aibridge_client_new(cProvider, cOpts) + if ptr == nil { + return nil, readLastError() + } + + client := &Client{ptr: ptr} + // RAII:GC 时若用户忘记 Close,兜底释放(best-effort,不应依赖) + runtime.SetFinalizer(client, func(c *Client) { + if c.ptr != nil { + C.aibridge_client_destroy(c.ptr) + c.ptr = nil + } + }) + return client, nil +} + +// Start 启动客户端(初始化适配器,对应 aibridge_client_start) +// +// 返回 0 成功;负数为错误码。 +func (c *Client) Start() error { + if c.ptr == nil { + return newFfiError("client 句柄为空(已 Close 或未初始化)") + } + // 同一 goroutine 调用 + 读取 last_error,保证线程局部语义 + status := C.aibridge_client_start(c.ptr) + if int32(status) != statusOK { + return readLastError() + } + return nil +} + +// Close 释放客户端句柄(对应 aibridge_client_destroy)。 +// +// 必须调用,否则内存泄漏。可多次调用(第二次 no-op)。 +func (c *Client) Close() { + if c.ptr != nil { + C.aibridge_client_destroy(c.ptr) + c.ptr = nil + runtime.SetFinalizer(c, nil) // 取消 finalizer,避免重复释放 + } +} + +// Chat 文本对话(阻塞,对应 aibridge_client_chat) +// +// 把 req 序列化为 JSON 传入 FFI,FFI 返回 ChatCompletion 的 JSON,反序列化为 Go struct。 +func (c *Client) Chat(req *ChatRequest) (*ChatCompletion, error) { + if c.ptr == nil { + return nil, newFfiError("client 句柄为空(已 Close 或未初始化)") + } + reqJSON, err := json.Marshal(req) + if err != nil { + return nil, newFfiError("ChatRequest JSON 序列化失败: " + err.Error()) + } + + cReq := C.CString(string(reqJSON)) + defer C.free(unsafe.Pointer(cReq)) + + var outResp *C.char + // aibridge_client_chat 内部 block_on(async),复杂 struct 走 JSON 边界 + status := C.aibridge_client_chat(c.ptr, cReq, &outResp) + if int32(status) != statusOK { + // 失败时 outResp 应为 nil,但稳妥起见仍检查释放 + if outResp != nil { + C.aibridge_string_free(outResp) + } + return nil, readLastError() + } + defer C.aibridge_string_free(outResp) // Rust 分配的 char* 必须释放 + + // C.GoString 拷贝到 Go 堆后即可安全使用 + respStr := C.GoString(outResp) + var completion ChatCompletion + if err := json.Unmarshal([]byte(respStr), &completion); err != nil { + return nil, newFfiError("ChatCompletion JSON 反序列化失败: " + err.Error()) + } + return &completion, nil +} + +// Speech 文字转语音(阻塞,对应 aibridge_client_speech) +// +// 二进制音频走 aibridge_bytes_t(避免 base64 膨胀), +// meta(SpeechResult 不含 audio_data)走 JSON。 +// 两者均为 Rust 分配,必须分别调 aibridge_bytes_free / aibridge_string_free。 +func (c *Client) Speech(req *SpeechRequest) (*SpeechResult, error) { + if c.ptr == nil { + return nil, newFfiError("client 句柄为空(已 Close 或未初始化)") + } + reqJSON, err := json.Marshal(req) + if err != nil { + return nil, newFfiError("SpeechRequest JSON 序列化失败: " + err.Error()) + } + + cReq := C.CString(string(reqJSON)) + defer C.free(unsafe.Pointer(cReq)) + + var outAudio *C.aibridge_bytes_t + var outMeta *C.char + + status := C.aibridge_client_speech(c.ptr, cReq, &outAudio, &outMeta) + if int32(status) != statusOK { + // 失败时仍可能分配了 audio,稳妥释放 + if outAudio != nil { + C.aibridge_bytes_free(outAudio) + } + if outMeta != nil { + C.aibridge_string_free(outMeta) + } + return nil, readLastError() + } + + // meta JSON 必释放 + if outMeta != nil { + defer C.aibridge_string_free(outMeta) + } + // audio bytes 必释放 + if outAudio != nil { + defer C.aibridge_bytes_free(outAudio) + } + + // 解析 meta JSON + result := &SpeechResult{} + if outMeta != nil { + metaStr := C.GoString(outMeta) + if err := json.Unmarshal([]byte(metaStr), result); err != nil { + return nil, newFfiError("SpeechResult meta JSON 反序列化失败: " + err.Error()) + } + } + + // 拷贝二进制音频数据到 Go 切片(必须在 bytes_free 之前完成拷贝) + if outAudio != nil { + // aibridge_bytes_t 布局:{ const uint8_t* ptr; uintptr_t len; } + // 用 unsafe 读取 ptr 和 len + audioBytes := (*C.aibridge_bytes_t)(unsafe.Pointer(outAudio)) + if audioBytes.ptr != nil && audioBytes.len > 0 { + // C.GoBytes 会拷贝数据,拷贝完成后即可释放 Rust 侧内存 + result.AudioData = C.GoBytes(unsafe.Pointer(audioBytes.ptr), C.int(int64(audioBytes.len))) + } + } + + return result, nil +} + +// readLastError 读取当前线程的 last_error 并转为 Go error +// +// 必须在 FFI 调用失败后立即在同一 goroutine 调用(线程局部语义)。 +// 返回的指针仅在当前线程的下一次 FFI 调用前有效,故此函数立即拷贝。 +func readLastError() error { + errPtr := C.aibridge_last_error() + if errPtr == nil { + return newFfiError("FFI 调用失败但 last_error 为空") + } + errJSON := C.GoString(errPtr) // 立即拷贝到 Go 堆 + return parseErrorJSON(errJSON) +} diff --git a/bindings/go/error.go b/bindings/go/error.go new file mode 100644 index 0000000..636d56f --- /dev/null +++ b/bindings/go/error.go @@ -0,0 +1,101 @@ +// Package aibridge - 错误处理 +// +// FFI 错误模型:aibridge_status_t 返回码(0 成功 / 负数错误类别)+ +// aibridge_last_error() 线程局部 JSON(含 code/message/details/retryable)。 +// +// 本文件把 FFI 错误映射为 Go error + 类型断言接口: +// +// var err error = ... +// if ae, ok := err.(AibridgeError); ok { +// fmt.Println(ae.Code()) // "rate_limit" 等 +// } +package aibridge + +import ( + "encoding/json" + "fmt" +) + +// AibridgeError 是 AIBridge 错误的类型断言接口 +// +// 对应设计文档 9.3 节:Go 用 error + 类型断言接口 type AibridgeError interface{ Code() string } +type AibridgeError interface { + error + Code() string // 错误码(如 "rate_limit"、"validation_error") + Retryable() bool // 是否可重试 + Message() string // 原始错误消息 +} + +// aibridgeErrorImpl 是 AibridgeError 的内部实现 +type aibridgeErrorImpl struct { + CodeStr string `json:"code"` + Msg string `json:"message"` + Details json.RawMessage `json:"details,omitempty"` + Retry bool `json:"retryable"` +} + +// Error 实现 error 接口 +func (e *aibridgeErrorImpl) Error() string { + if e.Retry { + return fmt.Sprintf("aibridge[%s] (retryable): %s", e.CodeStr, e.Msg) + } + return fmt.Sprintf("aibridge[%s]: %s", e.CodeStr, e.Msg) +} + +// Code 返回错误码 +func (e *aibridgeErrorImpl) Code() string { return e.CodeStr } + +// Retryable 返回是否可重试 +func (e *aibridgeErrorImpl) Retryable() bool { return e.Retry } + +// Message 返回原始错误消息 +func (e *aibridgeErrorImpl) Message() string { return e.Msg } + +// parseErrorJSON 把 last_error 的 JSON 字符串解析为 AibridgeError +// +// 输入格式:{"code":"...","message":"...","details":...,"retryable":bool} +// 若 JSON 解析失败,回退为通用 ffi_error。 +func parseErrorJSON(jsonStr string) AibridgeError { + if jsonStr == "" { + return &aibridgeErrorImpl{ + CodeStr: "ffi_error", + Msg: "未知错误(last_error 为空)", + } + } + var e aibridgeErrorImpl + if err := json.Unmarshal([]byte(jsonStr), &e); err != nil { + // JSON 解析失败,回退 + return &aibridgeErrorImpl{ + CodeStr: "ffi_error", + Msg: fmt.Sprintf("last_error JSON 解析失败: %s (原始: %s)", err, jsonStr), + } + } + if e.CodeStr == "" { + e.CodeStr = "ffi_error" + } + return &e +} + +// newFfiError 构造一个非业务类别的 FFI 层错误(如句柄为空、JSON 解析失败等) +func newFfiError(msg string) AibridgeError { + return &aibridgeErrorImpl{ + CodeStr: "ffi_error", + Msg: msg, + } +} + +// 常见错误码常量(与 aibridge.h 的 AIBRIDGE_ERR_* 宏对齐,供需要按码判断时使用) +const ( + errCodeAuthentication = "authentication" + errCodeRateLimit = "rate_limit" + errCodeValidation = "validation_error" + errCodeModelNotFound = "model_not_found" + errCodeAPI = "api_error" + errCodeNetwork = "network" + errCodeTimeout = "timeout" + errCodeUnsupportedCapability = "unsupported_capability" + errCodeProviderNotFound = "provider_not_found" + errCodeVoiceNotAvailable = "voice_not_available" + errCodeServiceUnavailable = "service_unavailable" + errCodeFFI = "ffi_error" +) diff --git a/bindings/go/example/hello.go b/bindings/go/example/hello.go new file mode 100644 index 0000000..7d00fee --- /dev/null +++ b/bindings/go/example/hello.go @@ -0,0 +1,140 @@ +// AIBridge Go 绑定 - Hello World 示例 +// +// 使用 echo adapter(provider="echo",免认证)验证跨语言管线: +// - Chat:回显最后一条 user 消息 + " [echo]" +// - ChatStream:3 个 chunk(role / 前半段 / 后半段+finish) +// - Speech:返回 15 字节 mock 音频(b"mock-audio-data") +// +// 运行: +// +// cd bindings/go +// CGO_ENABLED=1 go run ./example +// +// 需先 cargo build -p aibridge-ffi 产出 target/debug/libaibridge.dylib。 +// macOS 若报 dylib 找不到,设置: +// +// DYLD_LIBRARY_PATH=/Users/skywing/agn-sdk/target/debug +package main + +import ( + "fmt" + "os" + + aibridge "github.com/aibridge/aibridge-go" +) + +func main() { + fmt.Println("=== AIBridge Go 绑定 Hello World ===") + fmt.Println() + + // 1. 创建并启动 echo 客户端(免认证) + client, err := aibridge.NewClient("echo", nil) + if err != nil { + fmt.Fprintf(os.Stderr, "NewClient 失败: %v\n", err) + os.Exit(1) + } + defer client.Close() // RAII:确保释放句柄 + + if err := client.Start(); err != nil { + fmt.Fprintf(os.Stderr, "Start 失败: %v\n", err) + os.Exit(1) + } + fmt.Println("[OK] echo 客户端已创建并启动") + fmt.Println() + + // 2. Chat:回显 hello + " [echo]" + fmt.Println("--- Chat ---") + chatReq := &aibridge.ChatRequest{ + Model: "echo-chat", + Messages: []aibridge.ChatMessage{ + aibridge.NewUserTextMessage("hello"), + }, + } + completion, err := client.Chat(chatReq) + if err != nil { + fmt.Fprintf(os.Stderr, "Chat 失败: %v\n", err) + os.Exit(1) + } + content := "" + if len(completion.Choices) > 0 { + content = completion.Choices[0].Message.Content + } + fmt.Printf("[OK] Chat 返回: id=%s, model=%s\n", completion.ID, completion.Model) + fmt.Printf(" choices[0].message.content = %q (期望 %q)\n", content, "hello [echo]") + if content != "hello [echo]" { + fmt.Fprintf(os.Stderr, " [FAIL] 内容不匹配\n") + os.Exit(1) + } + fmt.Println() + + // 3. ChatStream:3 个 chunk + fmt.Println("--- ChatStream ---") + streamReq := &aibridge.ChatRequest{ + Model: "echo-chat", + Messages: []aibridge.ChatMessage{ + aibridge.NewUserTextMessage("hello"), + }, + } + stream, err := client.ChatStream(streamReq) + if err != nil { + fmt.Fprintf(os.Stderr, "ChatStream 失败: %v\n", err) + os.Exit(1) + } + + chunkCount := 0 + var assembledContent string + var finishReason string + for chunk := range stream.Ch() { + chunkCount++ + if len(chunk.Choices) > 0 { + delta := chunk.Choices[0].Delta + if delta.Role != "" { + fmt.Printf(" chunk %d: role=%q\n", chunkCount, delta.Role) + } + if delta.Content != "" { + fmt.Printf(" chunk %d: content=%q\n", chunkCount, delta.Content) + assembledContent += delta.Content + } + if chunk.Choices[0].FinishReason != "" { + finishReason = chunk.Choices[0].FinishReason + fmt.Printf(" chunk %d: finish_reason=%q\n", chunkCount, finishReason) + } + } + } + if err := stream.Err(); err != nil { + fmt.Fprintf(os.Stderr, " [FAIL] 流式错误: %v\n", err) + os.Exit(1) + } + fmt.Printf("[OK] ChatStream 收到 %d 个 chunk (期望 3)\n", chunkCount) + fmt.Printf(" 拼接内容 = %q (期望 %q)\n", assembledContent, "hello [echo]") + fmt.Printf(" finish_reason = %q\n", finishReason) + if chunkCount != 3 { + fmt.Fprintf(os.Stderr, " [FAIL] chunk 数量不匹配\n") + os.Exit(1) + } + fmt.Println() + + // 4. Speech:15 字节 mock 音频 + fmt.Println("--- Speech ---") + speechReq := &aibridge.SpeechRequest{ + Model: "echo-tts", + Input: "hello", + Voice: aibridge.SingleVoice("alloy"), + } + speech, err := client.Speech(speechReq) + if err != nil { + fmt.Fprintf(os.Stderr, "Speech 失败: %v\n", err) + os.Exit(1) + } + fmt.Printf("[OK] Speech 返回: model=%s, format=%s, content_type=%s\n", + speech.Model, speech.Format, speech.ContentType) + fmt.Printf(" len(AudioData) = %d (期望 15)\n", len(speech.AudioData)) + fmt.Printf(" AudioData = %q\n", string(speech.AudioData)) + if len(speech.AudioData) != 15 { + fmt.Fprintf(os.Stderr, " [FAIL] 音频字节数不匹配\n") + os.Exit(1) + } + fmt.Println() + + fmt.Println("=== 全部通过 ===") +} diff --git a/bindings/go/go.mod b/bindings/go/go.mod new file mode 100644 index 0000000..a972d8c --- /dev/null +++ b/bindings/go/go.mod @@ -0,0 +1,3 @@ +module github.com/aibridge/aibridge-go + +go 1.26 diff --git a/bindings/go/model.go b/bindings/go/model.go new file mode 100644 index 0000000..7029024 --- /dev/null +++ b/bindings/go/model.go @@ -0,0 +1,157 @@ +// Package aibridge - 数据模型 +// +// 本文件定义 AIBridge Go 绑定的数据模型(请求/响应 struct), +// 与 Rust 端 aibridge-core 的 serde struct 字段一一对应(JSON 边界)。 +// 对应:crates/aibridge-core/src/model/{chat,audio}.rs +package aibridge + +import "encoding/json" + +// ChatMessage 表示一条对话消息(与 Rust ChatMessage 对齐,role 作为 tag) +// +// JSON 序列化形如:{"role":"user","content":"hello"} +type ChatMessage struct { + Role string `json:"role"` // 角色:system / user / assistant / tool + Content json.RawMessage `json:"content"` // 内容:字符串或多模态部件列表(用 RawMessage 兼容两种形态) + Name string `json:"name,omitempty"` // 发送者名称(可选) + ToolCallID string `json:"tool_call_id,omitempty"` // 工具调用 ID(仅 role=tool) +} + +// NewUserTextMessage 构造一条纯文本用户消息 +func NewUserTextMessage(content string) ChatMessage { + // content 序列化为 JSON 字符串(带引号),与 Rust UserContent::Text 对齐 + b, _ := json.Marshal(content) + return ChatMessage{ + Role: "user", + Content: b, + } +} + +// NewSystemMessage 构造一条系统消息 +func NewSystemMessage(content string) ChatMessage { + return ChatMessage{ + Role: "system", + Content: json.RawMessage(`"` + jsonEscapeString(content) + `"`), + } +} + +// jsonEscapeString 对字符串做 JSON 转义(用于手工拼装带引号的 content) +func jsonEscapeString(s string) string { + b, _ := json.Marshal(s) + // 去掉首尾引号 + return string(b[1 : len(b)-1]) +} + +// ChatRequest 文本对话请求(对应 Rust ChatRequest) +type ChatRequest struct { + Model string `json:"model"` + Messages []ChatMessage `json:"messages"` + + Temperature float64 `json:"temperature,omitempty"` + TopP float64 `json:"top_p,omitempty"` + MaxTokens uint32 `json:"max_tokens,omitempty"` + N uint32 `json:"n,omitempty"` + PresencePenalty float64 `json:"presence_penalty,omitempty"` + FrequencyPenalty float64 `json:"frequency_penalty,omitempty"` + Seed uint64 `json:"seed,omitempty"` + Stream bool `json:"stream,omitempty"` + User string `json:"user,omitempty"` + + // 厂商特有参数透传 + Extra map[string]json.RawMessage `json:"extra,omitempty"` +} + +// ChatCompletion 文本对话完成结果(对应 Rust ChatCompletion) +type ChatCompletion struct { + ID string `json:"id"` + Object string `json:"object"` + Created uint64 `json:"created"` + Model string `json:"model"` + Choices []ChatChoice `json:"choices"` + Usage *ChatUsage `json:"usage,omitempty"` + ServiceTier string `json:"service_tier,omitempty"` + SystemFingerprint string `json:"system_fingerprint,omitempty"` +} + +// ChatChoice 对话选项 +type ChatChoice struct { + Index int `json:"index"` + Message ChoiceMessage `json:"message"` + FinishReason string `json:"finish_reason,omitempty"` +} + +// ChoiceMessage 完成结果中的消息 +type ChoiceMessage struct { + Role string `json:"role"` + Content string `json:"content,omitempty"` +} + +// ChatUsage Token 使用统计 +type ChatUsage struct { + PromptTokens uint64 `json:"prompt_tokens"` + CompletionTokens uint64 `json:"completion_tokens"` + TotalTokens uint64 `json:"total_tokens"` +} + +// ChatCompletionChunk 流式对话块(对应 Rust ChatCompletionChunk) +type ChatCompletionChunk struct { + ID string `json:"id"` + Object string `json:"object"` + Created uint64 `json:"created"` + Model string `json:"model"` + Choices []ChatCompletionDelta `json:"choices"` + Usage *ChatUsage `json:"usage,omitempty"` +} + +// ChatCompletionDelta 流式增量 +type ChatCompletionDelta struct { + Index int `json:"index"` + Delta DeltaMessage `json:"delta"` + FinishReason string `json:"finish_reason,omitempty"` +} + +// DeltaMessage 流式增量消息 +type DeltaMessage struct { + Role string `json:"role,omitempty"` + Content string `json:"content,omitempty"` +} + +// SpeechRequest 文字转语音请求(对应 Rust SpeechRequest) +// +// 注意:Voice 字段是 VoiceSpec 对象 {"voices":[...]},不是纯字符串。 +type SpeechRequest struct { + Model string `json:"model"` + Input string `json:"input"` + Voice VoiceSpec `json:"voice"` + ResponseFormat string `json:"response_format,omitempty"` + Speed float64 `json:"speed,omitempty"` + Volume float64 `json:"volume,omitempty"` + Pitch float64 `json:"pitch,omitempty"` + Emotion string `json:"emotion,omitempty"` + Style string `json:"style,omitempty"` + + Extra map[string]json.RawMessage `json:"extra,omitempty"` +} + +// VoiceSpec 音色规格(支持候选列表用于自动降级) +type VoiceSpec struct { + Voices []string `json:"voices"` +} + +// SingleVoice 构造单个音色的 VoiceSpec +func SingleVoice(v string) VoiceSpec { + return VoiceSpec{Voices: []string{v}} +} + +// SpeechResult 文字转语音结果(对应 Rust SpeechResult,audio_data 不参与序列化) +// +// 二进制音频数据通过 FFI 的 aibridge_bytes_t 单独传递,本结构体由 Speech() 填充。 +type SpeechResult struct { + AudioData []byte // 二进制音频数据(FFI 单独传递,不来自 JSON) + AudioURL string `json:"audio_url,omitempty"` + AudioBase64 string `json:"audio_base64,omitempty"` + ContentType string `json:"content_type"` + Format string `json:"format"` + Duration float64 `json:"duration,omitempty"` + Model string `json:"model,omitempty"` +} diff --git a/bindings/go/stream.go b/bindings/go/stream.go new file mode 100644 index 0000000..93d09f5 --- /dev/null +++ b/bindings/go/stream.go @@ -0,0 +1,152 @@ +// Package aibridge - 流式文本对话 +// +// 流式桥接策略(设计文档第 8 节): +// - aibridge_client_chat_stream 创建 stream 句柄 +// - goroutine 内串行循环 aibridge_stream_next(0=chunk/1=EOF/负=错) +// - chunk 反序列化为 ChatCompletionChunk,push 到 channel +// - EOF 或错误时关闭 channel,并 aibridge_stream_destroy 释放句柄 +// +// FFI 遗留:stream_next 串行(同一 stream 不可并发 next),goroutine 内单线程循环保证。 +// cancel:关闭 chan 或调用方放弃读取后,stream_destroy 仍会执行(defer),触发 Rust drop → tokio task abort。 +package aibridge + +/* +#cgo CFLAGS: -I${SRCDIR}/../../crates/aibridge-ffi/include + +#include +#include "aibridge.h" +*/ +import "C" + +import ( + "encoding/json" + "unsafe" +) + +// ChatStream 表示一个流式对话句柄,封装 stream 句柄与读取 channel。 +// +// 用法: +// +// ch, err := client.ChatStream(req) +// if err != nil { ... } +// for chunk := range ch { +// fmt.Println(chunk.Choices[0].Delta.Content) +// } +// // channel 关闭即表示流结束(正常 EOF 或错误) +// +// 错误获取:若流以错误结束,channel 关闭后可调 ch.Err() 获取错误。 +type ChatStream struct { + ptr *C.aibridge_stream_t + ch chan ChatCompletionChunk + errCh chan error + closed bool +} + +// ChatStream 创建流式对话(对应 aibridge_client_chat_stream) +// +// 返回的 ChatStream 已在后台 goroutine 中开始拉取 chunk, +// 调用方 for range chan 即可消费。流结束后 channel 自动关闭。 +// +// 注意:流式过程中调用方放弃读取不会泄漏——goroutine 会在下次 stream_next +// 返回后检测到 channel 阻塞并最终销毁 stream;但若调用方完全不读, +// goroutine 会阻塞在 channel 发送上。建议用 context 控制或显式消费。 +func (c *Client) ChatStream(req *ChatRequest) (*ChatStream, error) { + if c.ptr == nil { + return nil, newFfiError("client 句柄为空(已 Close 或未初始化)") + } + reqJSON, err := json.Marshal(req) + if err != nil { + return nil, newFfiError("ChatRequest JSON 序列化失败: " + err.Error()) + } + + cReq := C.CString(string(reqJSON)) + defer C.free(unsafe.Pointer(cReq)) + + var outStream *C.aibridge_stream_t + status := C.aibridge_client_chat_stream(c.ptr, cReq, &outStream) + if int32(status) != statusOK { + if outStream != nil { + C.aibridge_stream_destroy(outStream) + } + return nil, readLastError() + } + + cs := &ChatStream{ + ptr: outStream, + ch: make(chan ChatCompletionChunk), + errCh: make(chan error, 1), // 缓冲 1,避免 goroutine 因无人读取错误而泄漏 + } + + // 启动后台 goroutine 串行拉取 chunk + go cs.pullLoop() + + return cs, nil +} + +// pullLoop 后台 goroutine:串行调用 aibridge_stream_next +// +// - 0 (STREAM_CHUNK):拉到 chunk,反序列化后发送到 channel +// - 1 (STREAM_END):流正常结束,关闭 channel +// - 负数:错误,记录到 errCh,关闭 channel +// 无论何种结束,最后都 aibridge_stream_destroy 释放句柄。 +func (cs *ChatStream) pullLoop() { + defer close(cs.ch) + // 无论正常结束还是错误,最后都 aibridge_stream_destroy 释放句柄。 + // 注意:destroy 后不能再访问 cs.ptr,故用临时变量保存原始指针。 + streamPtr := cs.ptr + defer func() { + C.aibridge_stream_destroy(streamPtr) + cs.ptr = nil + }() + + for { + var outChunk *C.char + // 串行拉取(同一 stream 不可并发 next,本 goroutine 独占) + status := C.aibridge_stream_next(streamPtr, &outChunk) + + switch int32(status) { + case statusStreamChunk: + // 拉到一个 chunk + if outChunk == nil { + cs.errCh <- newFfiError("stream_next 返回 chunk 但 out_chunk_json 为空") + return + } + chunkStr := C.GoString(outChunk) + C.aibridge_string_free(outChunk) // chunk JSON 由 Rust 分配,必须释放 + + var chunk ChatCompletionChunk + if err := json.Unmarshal([]byte(chunkStr), &chunk); err != nil { + cs.errCh <- newFfiError("ChatCompletionChunk JSON 反序列化失败: " + err.Error()) + return + } + cs.ch <- chunk + + case statusStreamEnd: + // 流正常结束 + return + + default: + // 负数:错误 + cs.errCh <- readLastError() + return + } + } +} + +// Ch 返回流式 chunk 的只读 channel +func (cs *ChatStream) Ch() <-chan ChatCompletionChunk { + return cs.ch +} + +// Err 返回流式过程中的错误(若以错误结束) +// +// 必须在 channel 关闭后(for range 退出后)调用。 +// 若流正常结束,返回 nil。 +func (cs *ChatStream) Err() error { + select { + case err := <-cs.errCh: + return err + default: + return nil + } +} From 5118eebe4b4af165b1b7aead22bbd1cc8f3602a6 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 11:57:22 +0800 Subject: [PATCH 08/55] =?UTF-8?q?feat(aibridge-jvm):=20=E9=98=B6=E6=AE=B50?= =?UTF-8?q?.6=20JNA=20=E7=BB=91=E5=AE=9A=20+=20hello=20world?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- bindings/jvm/.gitignore | 16 + bindings/jvm/build.gradle.kts | 63 ++++ .../jvm/gradle/wrapper/gradle-wrapper.jar | Bin 0 -> 43583 bytes .../gradle/wrapper/gradle-wrapper.properties | 7 + bindings/jvm/gradlew | 252 +++++++++++++ bindings/jvm/gradlew.bat | 94 +++++ bindings/jvm/settings.gradle.kts | 1 + .../java/io/aibridge/AibridgeException.java | 88 +++++ .../main/java/io/aibridge/AibridgeNative.java | 179 ++++++++++ .../main/java/io/aibridge/ApiException.java | 8 + .../io/aibridge/AuthenticationException.java | 8 + .../src/main/java/io/aibridge/ChatChoice.java | 22 ++ .../main/java/io/aibridge/ChatCompletion.java | 38 ++ .../java/io/aibridge/ChatCompletionChunk.java | 34 ++ .../java/io/aibridge/ChatCompletionDelta.java | 19 + .../main/java/io/aibridge/ChatMessage.java | 50 +++ .../main/java/io/aibridge/ChatRequest.java | 73 ++++ .../src/main/java/io/aibridge/ChatStream.java | 265 ++++++++++++++ .../src/main/java/io/aibridge/ChatUsage.java | 21 ++ .../main/java/io/aibridge/ChoiceMessage.java | 24 ++ .../jvm/src/main/java/io/aibridge/Client.java | 330 ++++++++++++++++++ .../main/java/io/aibridge/DeltaMessage.java | 24 ++ .../jvm/src/main/java/io/aibridge/Hello.java | 133 +++++++ .../io/aibridge/ModelNotFoundException.java | 8 + .../java/io/aibridge/NetworkException.java | 8 + .../aibridge/ProviderNotFoundException.java | 8 + .../java/io/aibridge/RateLimitException.java | 12 + .../aibridge/ServiceUnavailableException.java | 8 + .../main/java/io/aibridge/SpeechRequest.java | 63 ++++ .../main/java/io/aibridge/SpeechResult.java | 34 ++ .../java/io/aibridge/SpeechResultFull.java | 36 ++ .../java/io/aibridge/TimeoutException.java | 8 + .../UnsupportedCapabilityException.java | 8 + .../java/io/aibridge/ValidationException.java | 8 + .../aibridge/VoiceNotAvailableException.java | 8 + .../src/main/java/io/aibridge/VoiceSpec.java | 42 +++ 36 files changed, 2000 insertions(+) create mode 100644 bindings/jvm/.gitignore create mode 100644 bindings/jvm/build.gradle.kts create mode 100644 bindings/jvm/gradle/wrapper/gradle-wrapper.jar create mode 100644 bindings/jvm/gradle/wrapper/gradle-wrapper.properties create mode 100755 bindings/jvm/gradlew create mode 100644 bindings/jvm/gradlew.bat create mode 100644 bindings/jvm/settings.gradle.kts create mode 100644 bindings/jvm/src/main/java/io/aibridge/AibridgeException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/AibridgeNative.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ApiException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/AuthenticationException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatChoice.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatCompletion.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatCompletionChunk.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatCompletionDelta.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatMessage.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatRequest.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatStream.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChatUsage.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ChoiceMessage.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/Client.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/DeltaMessage.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/Hello.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ModelNotFoundException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/NetworkException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ProviderNotFoundException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/RateLimitException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ServiceUnavailableException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/SpeechRequest.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/SpeechResult.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/SpeechResultFull.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/TimeoutException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/UnsupportedCapabilityException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/ValidationException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/VoiceNotAvailableException.java create mode 100644 bindings/jvm/src/main/java/io/aibridge/VoiceSpec.java diff --git a/bindings/jvm/.gitignore b/bindings/jvm/.gitignore new file mode 100644 index 0000000..88e3fc0 --- /dev/null +++ b/bindings/jvm/.gitignore @@ -0,0 +1,16 @@ +# Gradle +.gradle/ +build/ +!gradle/wrapper/gradle-wrapper.jar + +# IDE +.idea/ +*.iml +.vscode/ +.settings/ +.classpath +.project +bin/ + +# OS +.DS_Store diff --git a/bindings/jvm/build.gradle.kts b/bindings/jvm/build.gradle.kts new file mode 100644 index 0000000..5a6998c --- /dev/null +++ b/bindings/jvm/build.gradle.kts @@ -0,0 +1,63 @@ +// AIBridge JVM 绑定构建脚本 +// +// 通过 JNA 调用 aibridge-ffi 的 cdylib(libaibridge.dylib / .so / .dll)。 +// 依赖: +// - net.java.dev.jna:jna:纯 Java 调 C ABI +// - com.fasterxml.jackson.*:JSON 边界序列化/反序列化 +// +// 运行 hello world: +// ./gradlew run +// 需通过 -Djava.library.path 或环境变量 DYLD_LIBRARY_PATH / LD_LIBRARY_PATH +// 指向 libaibridge 所在目录(默认 target/debug)。 + +plugins { + application + java +} + +group = "io.aibridge" +version = "0.1.0" + +java { + toolchain { + languageVersion = JavaLanguageVersion.of(21) + } +} + +repositories { + mavenCentral() +} + +dependencies { + // JNA:纯 Java 调 C ABI,无需 native Rust 绑定 + implementation("net.java.dev.jna:jna:5.15.0") + // Jackson:JSON 边界(反)序列化 + implementation("com.fasterxml.jackson.core:jackson-databind:2.18.1") + // JSR305 注解(@Nullable 等),提升 JNA Pointer 语义可读性 + implementation("com.google.code.findbugs:jsr305:3.0.2") +} + +application { + // Hello world 入口 + mainClass = "io.aibridge.Hello" + // 默认把 cargo debug 输出目录加入 java.library.path(可被命令行覆盖) + applicationDefaultJvmArgs = listOf( + "-Djava.library.path=${rootProject.projectDir}/../../target/debug", + "-Djna.library.path=${rootProject.projectDir}/../../target/debug" + ) +} + +tasks.withType { + options.encoding = "UTF-8" + // 保留参数名(便于反射/Jackson) + options.compilerArgs.addAll(listOf("-parameters")) +} + +tasks.withType { + // 传递系统属性(如 -D...),便于调试 + systemProperties = System.getProperties() as Map +} + +tasks.test { + useJUnitPlatform() +} diff --git a/bindings/jvm/gradle/wrapper/gradle-wrapper.jar b/bindings/jvm/gradle/wrapper/gradle-wrapper.jar new file mode 100644 index 0000000000000000000000000000000000000000..a4b76b9530d66f5e68d973ea569d8e19de379189 GIT binary patch literal 43583 zcma&N1CXTcmMvW9vTb(Rwr$&4wr$(C?dmSu>@vG-+vuvg^_??!{yS%8zW-#zn-LkA z5&1^$^{lnmUON?}LBF8_K|(?T0Ra(xUH{($5eN!MR#ZihR#HxkUPe+_R8Cn`RRs(P z_^*#_XlXmGv7!4;*Y%p4nw?{bNp@UZHv1?Um8r6)Fei3p@ClJn0ECfg1hkeuUU@Or zDaPa;U3fE=3L}DooL;8f;P0ipPt0Z~9P0)lbStMS)ag54=uL9ia-Lm3nh|@(Y?B`; zx_#arJIpXH!U{fbCbI^17}6Ri*H<>OLR%c|^mh8+)*h~K8Z!9)DPf zR2h?lbDZQ`p9P;&DQ4F0sur@TMa!Y}S8irn(%d-gi0*WxxCSk*A?3lGh=gcYN?FGl z7D=Js!i~0=u3rox^eO3i@$0=n{K1lPNU zwmfjRVmLOCRfe=seV&P*1Iq=^i`502keY8Uy-WNPwVNNtJFx?IwAyRPZo2Wo1+S(xF37LJZ~%i)kpFQ3Fw=mXfd@>%+)RpYQLnr}B~~zoof(JVm^^&f zxKV^+3D3$A1G;qh4gPVjhrC8e(VYUHv#dy^)(RoUFM?o%W-EHxufuWf(l*@-l+7vt z=l`qmR56K~F|v<^Pd*p~1_y^P0P^aPC##d8+HqX4IR1gu+7w#~TBFphJxF)T$2WEa zxa?H&6=Qe7d(#tha?_1uQys2KtHQ{)Qco)qwGjrdNL7thd^G5i8Os)CHqc>iOidS} z%nFEDdm=GXBw=yXe1W-ShHHFb?Cc70+$W~z_+}nAoHFYI1MV1wZegw*0y^tC*s%3h zhD3tN8b=Gv&rj}!SUM6|ajSPp*58KR7MPpI{oAJCtY~JECm)*m_x>AZEu>DFgUcby z1Qaw8lU4jZpQ_$;*7RME+gq1KySGG#Wql>aL~k9tLrSO()LWn*q&YxHEuzmwd1?aAtI zBJ>P=&$=l1efe1CDU;`Fd+_;&wI07?V0aAIgc(!{a z0Jg6Y=inXc3^n!U0Atk`iCFIQooHqcWhO(qrieUOW8X(x?(RD}iYDLMjSwffH2~tB z)oDgNBLB^AJBM1M^c5HdRx6fBfka`(LD-qrlh5jqH~);#nw|iyp)()xVYak3;Ybik z0j`(+69aK*B>)e_p%=wu8XC&9e{AO4c~O1U`5X9}?0mrd*m$_EUek{R?DNSh(=br# z#Q61gBzEpmy`$pA*6!87 zSDD+=@fTY7<4A?GLqpA?Pb2z$pbCc4B4zL{BeZ?F-8`s$?>*lXXtn*NC61>|*w7J* z$?!iB{6R-0=KFmyp1nnEmLsA-H0a6l+1uaH^g%c(p{iT&YFrbQ$&PRb8Up#X3@Zsk zD^^&LK~111%cqlP%!_gFNa^dTYT?rhkGl}5=fL{a`UViaXWI$k-UcHJwmaH1s=S$4 z%4)PdWJX;hh5UoK?6aWoyLxX&NhNRqKam7tcOkLh{%j3K^4Mgx1@i|Pi&}<^5>hs5 zm8?uOS>%)NzT(%PjVPGa?X%`N2TQCKbeH2l;cTnHiHppPSJ<7y-yEIiC!P*ikl&!B z%+?>VttCOQM@ShFguHVjxX^?mHX^hSaO_;pnyh^v9EumqSZTi+#f&_Vaija0Q-e*| z7ulQj6Fs*bbmsWp{`auM04gGwsYYdNNZcg|ph0OgD>7O}Asn7^Z=eI>`$2*v78;sj-}oMoEj&@)9+ycEOo92xSyY344^ z11Hb8^kdOvbf^GNAK++bYioknrpdN>+u8R?JxG=!2Kd9r=YWCOJYXYuM0cOq^FhEd zBg2puKy__7VT3-r*dG4c62Wgxi52EMCQ`bKgf*#*ou(D4-ZN$+mg&7$u!! z-^+Z%;-3IDwqZ|K=ah85OLwkO zKxNBh+4QHh)u9D?MFtpbl)us}9+V!D%w9jfAMYEb>%$A;u)rrI zuBudh;5PN}_6J_}l55P3l_)&RMlH{m!)ai-i$g)&*M`eN$XQMw{v^r@-125^RRCF0 z^2>|DxhQw(mtNEI2Kj(;KblC7x=JlK$@78`O~>V!`|1Lm-^JR$-5pUANAnb(5}B}JGjBsliK4& zk6y(;$e&h)lh2)L=bvZKbvh@>vLlreBdH8No2>$#%_Wp1U0N7Ank!6$dFSi#xzh|( zRi{Uw%-4W!{IXZ)fWx@XX6;&(m_F%c6~X8hx=BN1&q}*( zoaNjWabE{oUPb!Bt$eyd#$5j9rItB-h*5JiNi(v^e|XKAj*8(k<5-2$&ZBR5fF|JA z9&m4fbzNQnAU}r8ab>fFV%J0z5awe#UZ|bz?Ur)U9bCIKWEzi2%A+5CLqh?}K4JHi z4vtM;+uPsVz{Lfr;78W78gC;z*yTch~4YkLr&m-7%-xc ztw6Mh2d>_iO*$Rd8(-Cr1_V8EO1f*^@wRoSozS) zy1UoC@pruAaC8Z_7~_w4Q6n*&B0AjOmMWa;sIav&gu z|J5&|{=a@vR!~k-OjKEgPFCzcJ>#A1uL&7xTDn;{XBdeM}V=l3B8fE1--DHjSaxoSjNKEM9|U9#m2<3>n{Iuo`r3UZp;>GkT2YBNAh|b z^jTq-hJp(ebZh#Lk8hVBP%qXwv-@vbvoREX$TqRGTgEi$%_F9tZES@z8Bx}$#5eeG zk^UsLBH{bc2VBW)*EdS({yw=?qmevwi?BL6*=12k9zM5gJv1>y#ML4!)iiPzVaH9% zgSImetD@dam~e>{LvVh!phhzpW+iFvWpGT#CVE5TQ40n%F|p(sP5mXxna+Ev7PDwA zamaV4m*^~*xV+&p;W749xhb_X=$|LD;FHuB&JL5?*Y2-oIT(wYY2;73<^#46S~Gx| z^cez%V7x$81}UWqS13Gz80379Rj;6~WdiXWOSsdmzY39L;Hg3MH43o*y8ibNBBH`(av4|u;YPq%{R;IuYow<+GEsf@R?=@tT@!}?#>zIIn0CoyV!hq3mw zHj>OOjfJM3F{RG#6ujzo?y32m^tgSXf@v=J$ELdJ+=5j|=F-~hP$G&}tDZsZE?5rX ztGj`!S>)CFmdkccxM9eGIcGnS2AfK#gXwj%esuIBNJQP1WV~b~+D7PJTmWGTSDrR` zEAu4B8l>NPuhsk5a`rReSya2nfV1EK01+G!x8aBdTs3Io$u5!6n6KX%uv@DxAp3F@{4UYg4SWJtQ-W~0MDb|j-$lwVn znAm*Pl!?Ps&3wO=R115RWKb*JKoexo*)uhhHBncEDMSVa_PyA>k{Zm2(wMQ(5NM3# z)jkza|GoWEQo4^s*wE(gHz?Xsg4`}HUAcs42cM1-qq_=+=!Gk^y710j=66(cSWqUe zklbm8+zB_syQv5A2rj!Vbw8;|$@C!vfNmNV!yJIWDQ>{+2x zKjuFX`~~HKG~^6h5FntRpnnHt=D&rq0>IJ9#F0eM)Y-)GpRjiN7gkA8wvnG#K=q{q z9dBn8_~wm4J<3J_vl|9H{7q6u2A!cW{bp#r*-f{gOV^e=8S{nc1DxMHFwuM$;aVI^ zz6A*}m8N-&x8;aunp1w7_vtB*pa+OYBw=TMc6QK=mbA-|Cf* zvyh8D4LRJImooUaSb7t*fVfih<97Gf@VE0|z>NcBwBQze);Rh!k3K_sfunToZY;f2 z^HmC4KjHRVg+eKYj;PRN^|E0>Gj_zagfRbrki68I^#~6-HaHg3BUW%+clM1xQEdPYt_g<2K+z!$>*$9nQ>; zf9Bei{?zY^-e{q_*|W#2rJG`2fy@{%6u0i_VEWTq$*(ZN37|8lFFFt)nCG({r!q#9 z5VK_kkSJ3?zOH)OezMT{!YkCuSSn!K#-Rhl$uUM(bq*jY? zi1xbMVthJ`E>d>(f3)~fozjg^@eheMF6<)I`oeJYx4*+M&%c9VArn(OM-wp%M<-`x z7sLP1&3^%Nld9Dhm@$3f2}87!quhI@nwd@3~fZl_3LYW-B?Ia>ui`ELg z&Qfe!7m6ze=mZ`Ia9$z|ARSw|IdMpooY4YiPN8K z4B(ts3p%2i(Td=tgEHX z0UQ_>URBtG+-?0E;E7Ld^dyZ;jjw0}XZ(}-QzC6+NN=40oDb2^v!L1g9xRvE#@IBR zO!b-2N7wVfLV;mhEaXQ9XAU+>=XVA6f&T4Z-@AX!leJ8obP^P^wP0aICND?~w&NykJ#54x3_@r7IDMdRNy4Hh;h*!u(Ol(#0bJdwEo$5437-UBjQ+j=Ic>Q2z` zJNDf0yO6@mr6y1#n3)s(W|$iE_i8r@Gd@!DWDqZ7J&~gAm1#~maIGJ1sls^gxL9LLG_NhU!pTGty!TbhzQnu)I*S^54U6Yu%ZeCg`R>Q zhBv$n5j0v%O_j{QYWG!R9W?5_b&67KB$t}&e2LdMvd(PxN6Ir!H4>PNlerpBL>Zvyy!yw z-SOo8caEpDt(}|gKPBd$qND5#a5nju^O>V&;f890?yEOfkSG^HQVmEbM3Ugzu+UtH zC(INPDdraBN?P%kE;*Ae%Wto&sgw(crfZ#Qy(<4nk;S|hD3j{IQRI6Yq|f^basLY; z-HB&Je%Gg}Jt@={_C{L$!RM;$$|iD6vu#3w?v?*;&()uB|I-XqEKqZPS!reW9JkLewLb!70T7n`i!gNtb1%vN- zySZj{8-1>6E%H&=V}LM#xmt`J3XQoaD|@XygXjdZ1+P77-=;=eYpoEQ01B@L*a(uW zrZeZz?HJsw_4g0vhUgkg@VF8<-X$B8pOqCuWAl28uB|@r`19DTUQQsb^pfqB6QtiT z*`_UZ`fT}vtUY#%sq2{rchyfu*pCg;uec2$-$N_xgjZcoumE5vSI{+s@iLWoz^Mf; zuI8kDP{!XY6OP~q5}%1&L}CtfH^N<3o4L@J@zg1-mt{9L`s^z$Vgb|mr{@WiwAqKg zp#t-lhrU>F8o0s1q_9y`gQNf~Vb!F%70f}$>i7o4ho$`uciNf=xgJ>&!gSt0g;M>*x4-`U)ysFW&Vs^Vk6m%?iuWU+o&m(2Jm26Y(3%TL; zA7T)BP{WS!&xmxNw%J=$MPfn(9*^*TV;$JwRy8Zl*yUZi8jWYF>==j~&S|Xinsb%c z2?B+kpet*muEW7@AzjBA^wAJBY8i|#C{WtO_or&Nj2{=6JTTX05}|H>N2B|Wf!*3_ z7hW*j6p3TvpghEc6-wufFiY!%-GvOx*bZrhZu+7?iSrZL5q9}igiF^*R3%DE4aCHZ zqu>xS8LkW+Auv%z-<1Xs92u23R$nk@Pk}MU5!gT|c7vGlEA%G^2th&Q*zfg%-D^=f z&J_}jskj|Q;73NP4<4k*Y%pXPU2Thoqr+5uH1yEYM|VtBPW6lXaetokD0u z9qVek6Q&wk)tFbQ8(^HGf3Wp16gKmr>G;#G(HRBx?F`9AIRboK+;OfHaLJ(P>IP0w zyTbTkx_THEOs%Q&aPrxbZrJlio+hCC_HK<4%f3ZoSAyG7Dn`=X=&h@m*|UYO-4Hq0 z-Bq&+Ie!S##4A6OGoC~>ZW`Y5J)*ouaFl_e9GA*VSL!O_@xGiBw!AF}1{tB)z(w%c zS1Hmrb9OC8>0a_$BzeiN?rkPLc9%&;1CZW*4}CDDNr2gcl_3z+WC15&H1Zc2{o~i) z)LLW=WQ{?ricmC`G1GfJ0Yp4Dy~Ba;j6ZV4r{8xRs`13{dD!xXmr^Aga|C=iSmor% z8hi|pTXH)5Yf&v~exp3o+sY4B^^b*eYkkCYl*T{*=-0HniSA_1F53eCb{x~1k3*`W zr~};p1A`k{1DV9=UPnLDgz{aJH=-LQo<5%+Em!DNN252xwIf*wF_zS^!(XSm(9eoj z=*dXG&n0>)_)N5oc6v!>-bd(2ragD8O=M|wGW z!xJQS<)u70m&6OmrF0WSsr@I%T*c#Qo#Ha4d3COcX+9}hM5!7JIGF>7<~C(Ear^Sn zm^ZFkV6~Ula6+8S?oOROOA6$C&q&dp`>oR-2Ym3(HT@O7Sd5c~+kjrmM)YmgPH*tL zX+znN>`tv;5eOfX?h{AuX^LK~V#gPCu=)Tigtq9&?7Xh$qN|%A$?V*v=&-2F$zTUv z`C#WyIrChS5|Kgm_GeudCFf;)!WH7FI60j^0o#65o6`w*S7R@)88n$1nrgU(oU0M9 zx+EuMkC>(4j1;m6NoGqEkpJYJ?vc|B zOlwT3t&UgL!pX_P*6g36`ZXQ; z9~Cv}ANFnJGp(;ZhS(@FT;3e)0)Kp;h^x;$*xZn*k0U6-&FwI=uOGaODdrsp-!K$Ac32^c{+FhI-HkYd5v=`PGsg%6I`4d9Jy)uW0y%) zm&j^9WBAp*P8#kGJUhB!L?a%h$hJgQrx!6KCB_TRo%9{t0J7KW8!o1B!NC)VGLM5! zpZy5Jc{`r{1e(jd%jsG7k%I+m#CGS*BPA65ZVW~fLYw0dA-H_}O zrkGFL&P1PG9p2(%QiEWm6x;U-U&I#;Em$nx-_I^wtgw3xUPVVu zqSuKnx&dIT-XT+T10p;yjo1Y)z(x1fb8Dzfn8e yu?e%!_ptzGB|8GrCfu%p?(_ zQccdaaVK$5bz;*rnyK{_SQYM>;aES6Qs^lj9lEs6_J+%nIiuQC*fN;z8md>r_~Mfl zU%p5Dt_YT>gQqfr@`cR!$NWr~+`CZb%dn;WtzrAOI>P_JtsB76PYe*<%H(y>qx-`Kq!X_; z<{RpAqYhE=L1r*M)gNF3B8r(<%8mo*SR2hu zccLRZwGARt)Hlo1euqTyM>^!HK*!Q2P;4UYrysje@;(<|$&%vQekbn|0Ruu_Io(w4#%p6ld2Yp7tlA`Y$cciThP zKzNGIMPXX%&Ud0uQh!uQZz|FB`4KGD?3!ND?wQt6!n*f4EmCoJUh&b?;B{|lxs#F- z31~HQ`SF4x$&v00@(P+j1pAaj5!s`)b2RDBp*PB=2IB>oBF!*6vwr7Dp%zpAx*dPr zb@Zjq^XjN?O4QcZ*O+8>)|HlrR>oD*?WQl5ri3R#2?*W6iJ>>kH%KnnME&TT@ZzrHS$Q%LC?n|e>V+D+8D zYc4)QddFz7I8#}y#Wj6>4P%34dZH~OUDb?uP%-E zwjXM(?Sg~1!|wI(RVuxbu)-rH+O=igSho_pDCw(c6b=P zKk4ATlB?bj9+HHlh<_!&z0rx13K3ZrAR8W)!@Y}o`?a*JJsD+twZIv`W)@Y?Amu_u zz``@-e2X}27$i(2=9rvIu5uTUOVhzwu%mNazS|lZb&PT;XE2|B&W1>=B58#*!~D&) zfVmJGg8UdP*fx(>Cj^?yS^zH#o-$Q-*$SnK(ZVFkw+er=>N^7!)FtP3y~Xxnu^nzY zikgB>Nj0%;WOltWIob|}%lo?_C7<``a5hEkx&1ku$|)i>Rh6@3h*`slY=9U}(Ql_< zaNG*J8vb&@zpdhAvv`?{=zDedJ23TD&Zg__snRAH4eh~^oawdYi6A3w8<Ozh@Kw)#bdktM^GVb zrG08?0bG?|NG+w^&JvD*7LAbjED{_Zkc`3H!My>0u5Q}m!+6VokMLXxl`Mkd=g&Xx z-a>m*#G3SLlhbKB!)tnzfWOBV;u;ftU}S!NdD5+YtOjLg?X}dl>7m^gOpihrf1;PY zvll&>dIuUGs{Qnd- zwIR3oIrct8Va^Tm0t#(bJD7c$Z7DO9*7NnRZorrSm`b`cxz>OIC;jSE3DO8`hX955ui`s%||YQtt2 z5DNA&pG-V+4oI2s*x^>-$6J?p=I>C|9wZF8z;VjR??Icg?1w2v5Me+FgAeGGa8(3S z4vg*$>zC-WIVZtJ7}o9{D-7d>zCe|z#<9>CFve-OPAYsneTb^JH!Enaza#j}^mXy1 z+ULn^10+rWLF6j2>Ya@@Kq?26>AqK{A_| zQKb*~F1>sE*=d?A?W7N2j?L09_7n+HGi{VY;MoTGr_)G9)ot$p!-UY5zZ2Xtbm=t z@dpPSGwgH=QtIcEulQNI>S-#ifbnO5EWkI;$A|pxJd885oM+ zGZ0_0gDvG8q2xebj+fbCHYfAXuZStH2j~|d^sBAzo46(K8n59+T6rzBwK)^rfPT+B zyIFw)9YC-V^rhtK`!3jrhmW-sTmM+tPH+;nwjL#-SjQPUZ53L@A>y*rt(#M(qsiB2 zx6B)dI}6Wlsw%bJ8h|(lhkJVogQZA&n{?Vgs6gNSXzuZpEyu*xySy8ro07QZ7Vk1!3tJphN_5V7qOiyK8p z#@jcDD8nmtYi1^l8ml;AF<#IPK?!pqf9D4moYk>d99Im}Jtwj6c#+A;f)CQ*f-hZ< z=p_T86jog%!p)D&5g9taSwYi&eP z#JuEK%+NULWus;0w32-SYFku#i}d~+{Pkho&^{;RxzP&0!RCm3-9K6`>KZpnzS6?L z^H^V*s!8<>x8bomvD%rh>Zp3>Db%kyin;qtl+jAv8Oo~1g~mqGAC&Qi_wy|xEt2iz zWAJEfTV%cl2Cs<1L&DLRVVH05EDq`pH7Oh7sR`NNkL%wi}8n>IXcO40hp+J+sC!W?!krJf!GJNE8uj zg-y~Ns-<~D?yqbzVRB}G>0A^f0!^N7l=$m0OdZuqAOQqLc zX?AEGr1Ht+inZ-Qiwnl@Z0qukd__a!C*CKuGdy5#nD7VUBM^6OCpxCa2A(X;e0&V4 zM&WR8+wErQ7UIc6LY~Q9x%Sn*Tn>>P`^t&idaOEnOd(Ufw#>NoR^1QdhJ8s`h^|R_ zXX`c5*O~Xdvh%q;7L!_!ohf$NfEBmCde|#uVZvEo>OfEq%+Ns7&_f$OR9xsihRpBb z+cjk8LyDm@U{YN>+r46?nn{7Gh(;WhFw6GAxtcKD+YWV?uge>;+q#Xx4!GpRkVZYu zzsF}1)7$?%s9g9CH=Zs+B%M_)+~*j3L0&Q9u7!|+T`^O{xE6qvAP?XWv9_MrZKdo& z%IyU)$Q95AB4!#hT!_dA>4e@zjOBD*Y=XjtMm)V|+IXzjuM;(l+8aA5#Kaz_$rR6! zj>#&^DidYD$nUY(D$mH`9eb|dtV0b{S>H6FBfq>t5`;OxA4Nn{J(+XihF(stSche7$es&~N$epi&PDM_N`As;*9D^L==2Q7Z2zD+CiU(|+-kL*VG+&9!Yb3LgPy?A zm7Z&^qRG_JIxK7-FBzZI3Q<;{`DIxtc48k> zc|0dmX;Z=W$+)qE)~`yn6MdoJ4co;%!`ddy+FV538Y)j(vg}5*k(WK)KWZ3WaOG!8 z!syGn=s{H$odtpqFrT#JGM*utN7B((abXnpDM6w56nhw}OY}0TiTG1#f*VFZr+^-g zbP10`$LPq_;PvrA1XXlyx2uM^mrjTzX}w{yuLo-cOClE8MMk47T25G8M!9Z5ypOSV zAJUBGEg5L2fY)ZGJb^E34R2zJ?}Vf>{~gB!8=5Z) z9y$>5c)=;o0HeHHSuE4U)#vG&KF|I%-cF6f$~pdYJWk_dD}iOA>iA$O$+4%@>JU08 zS`ep)$XLPJ+n0_i@PkF#ri6T8?ZeAot$6JIYHm&P6EB=BiaNY|aA$W0I+nz*zkz_z zkEru!tj!QUffq%)8y0y`T&`fuus-1p>=^hnBiBqD^hXrPs`PY9tU3m0np~rISY09> z`P3s=-kt_cYcxWd{de@}TwSqg*xVhp;E9zCsnXo6z z?f&Sv^U7n4`xr=mXle94HzOdN!2kB~4=%)u&N!+2;z6UYKUDqi-s6AZ!haB;@&B`? z_TRX0%@suz^TRdCb?!vNJYPY8L_}&07uySH9%W^Tc&1pia6y1q#?*Drf}GjGbPjBS zbOPcUY#*$3sL2x4v_i*Y=N7E$mR}J%|GUI(>WEr+28+V z%v5{#e!UF*6~G&%;l*q*$V?&r$Pp^sE^i-0$+RH3ERUUdQ0>rAq2(2QAbG}$y{de( z>{qD~GGuOk559Y@%$?N^1ApVL_a704>8OD%8Y%8B;FCt%AoPu8*D1 zLB5X>b}Syz81pn;xnB}%0FnwazlWfUV)Z-~rZg6~b z6!9J$EcE&sEbzcy?CI~=boWA&eeIa%z(7SE^qgVLz??1Vbc1*aRvc%Mri)AJaAG!p z$X!_9Ds;Zz)f+;%s&dRcJt2==P{^j3bf0M=nJd&xwUGlUFn?H=2W(*2I2Gdu zv!gYCwM10aeus)`RIZSrCK=&oKaO_Ry~D1B5!y0R=%!i2*KfXGYX&gNv_u+n9wiR5 z*e$Zjju&ODRW3phN925%S(jL+bCHv6rZtc?!*`1TyYXT6%Ju=|X;6D@lq$8T zW{Y|e39ioPez(pBH%k)HzFITXHvnD6hw^lIoUMA;qAJ^CU?top1fo@s7xT13Fvn1H z6JWa-6+FJF#x>~+A;D~;VDs26>^oH0EI`IYT2iagy23?nyJ==i{g4%HrAf1-*v zK1)~@&(KkwR7TL}L(A@C_S0G;-GMDy=MJn2$FP5s<%wC)4jC5PXoxrQBFZ_k0P{{s@sz+gX`-!=T8rcB(=7vW}^K6oLWMmp(rwDh}b zwaGGd>yEy6fHv%jM$yJXo5oMAQ>c9j`**}F?MCry;T@47@r?&sKHgVe$MCqk#Z_3S z1GZI~nOEN*P~+UaFGnj{{Jo@16`(qVNtbU>O0Hf57-P>x8Jikp=`s8xWs^dAJ9lCQ z)GFm+=OV%AMVqVATtN@|vp61VVAHRn87}%PC^RAzJ%JngmZTasWBAWsoAqBU+8L8u z4A&Pe?fmTm0?mK-BL9t+{y7o(7jm+RpOhL9KnY#E&qu^}B6=K_dB}*VlSEiC9fn)+V=J;OnN)Ta5v66ic1rG+dGAJ1 z1%Zb_+!$=tQ~lxQrzv3x#CPb?CekEkA}0MYSgx$Jdd}q8+R=ma$|&1a#)TQ=l$1tQ z=tL9&_^vJ)Pk}EDO-va`UCT1m#Uty1{v^A3P~83_#v^ozH}6*9mIjIr;t3Uv%@VeW zGL6(CwCUp)Jq%G0bIG%?{_*Y#5IHf*5M@wPo6A{$Um++Co$wLC=J1aoG93&T7Ho}P z=mGEPP7GbvoG!uD$k(H3A$Z))+i{Hy?QHdk>3xSBXR0j!11O^mEe9RHmw!pvzv?Ua~2_l2Yh~_!s1qS`|0~0)YsbHSz8!mG)WiJE| z2f($6TQtt6L_f~ApQYQKSb=`053LgrQq7G@98#igV>y#i==-nEjQ!XNu9 z~;mE+gtj4IDDNQJ~JVk5Ux6&LCSFL!y=>79kE9=V}J7tD==Ga+IW zX)r7>VZ9dY=V&}DR))xUoV!u(Z|%3ciQi_2jl}3=$Agc(`RPb z8kEBpvY>1FGQ9W$n>Cq=DIpski};nE)`p3IUw1Oz0|wxll^)4dq3;CCY@RyJgFgc# zKouFh!`?Xuo{IMz^xi-h=StCis_M7yq$u) z?XHvw*HP0VgR+KR6wI)jEMX|ssqYvSf*_3W8zVTQzD?3>H!#>InzpSO)@SC8q*ii- z%%h}_#0{4JG;Jm`4zg};BPTGkYamx$Xo#O~lBirRY)q=5M45n{GCfV7h9qwyu1NxOMoP4)jjZMxmT|IQQh0U7C$EbnMN<3)Kk?fFHYq$d|ICu>KbY_hO zTZM+uKHe(cIZfEqyzyYSUBZa8;Fcut-GN!HSA9ius`ltNebF46ZX_BbZNU}}ZOm{M2&nANL9@0qvih15(|`S~z}m&h!u4x~(%MAO$jHRWNfuxWF#B)E&g3ghSQ9|> z(MFaLQj)NE0lowyjvg8z0#m6FIuKE9lDO~Glg}nSb7`~^&#(Lw{}GVOS>U)m8bF}x zVjbXljBm34Cs-yM6TVusr+3kYFjr28STT3g056y3cH5Tmge~ASxBj z%|yb>$eF;WgrcOZf569sDZOVwoo%8>XO>XQOX1OyN9I-SQgrm;U;+#3OI(zrWyow3 zk==|{lt2xrQ%FIXOTejR>;wv(Pb8u8}BUpx?yd(Abh6? zsoO3VYWkeLnF43&@*#MQ9-i-d0t*xN-UEyNKeyNMHw|A(k(_6QKO=nKMCxD(W(Yop zsRQ)QeL4X3Lxp^L%wzi2-WVSsf61dqliPUM7srDB?Wm6Lzn0&{*}|IsKQW;02(Y&| zaTKv|`U(pSzuvR6Rduu$wzK_W-Y-7>7s?G$)U}&uK;<>vU}^^ns@Z!p+9?St1s)dG zK%y6xkPyyS1$~&6v{kl?Md6gwM|>mt6Upm>oa8RLD^8T{0?HC!Z>;(Bob7el(DV6x zi`I)$&E&ngwFS@bi4^xFLAn`=fzTC;aimE^!cMI2n@Vo%Ae-ne`RF((&5y6xsjjAZ zVguVoQ?Z9uk$2ON;ersE%PU*xGO@T*;j1BO5#TuZKEf(mB7|g7pcEA=nYJ{s3vlbg zd4-DUlD{*6o%Gc^N!Nptgay>j6E5;3psI+C3Q!1ZIbeCubW%w4pq9)MSDyB{HLm|k zxv-{$$A*pS@csolri$Ge<4VZ}e~78JOL-EVyrbxKra^d{?|NnPp86!q>t<&IP07?Z z^>~IK^k#OEKgRH+LjllZXk7iA>2cfH6+(e&9ku5poo~6y{GC5>(bRK7hwjiurqAiZ zg*DmtgY}v83IjE&AbiWgMyFbaRUPZ{lYiz$U^&Zt2YjG<%m((&_JUbZcfJ22(>bi5 z!J?<7AySj0JZ&<-qXX;mcV!f~>G=sB0KnjWca4}vrtunD^1TrpfeS^4dvFr!65knK zZh`d;*VOkPs4*-9kL>$GP0`(M!j~B;#x?Ba~&s6CopvO86oM?-? zOw#dIRc;6A6T?B`Qp%^<U5 z19x(ywSH$_N+Io!6;e?`tWaM$`=Db!gzx|lQ${DG!zb1Zl&|{kX0y6xvO1o z220r<-oaS^^R2pEyY;=Qllqpmue|5yI~D|iI!IGt@iod{Opz@*ml^w2bNs)p`M(Io z|E;;m*Xpjd9l)4G#KaWfV(t8YUn@A;nK^#xgv=LtnArX|vWQVuw3}B${h+frU2>9^ z!l6)!Uo4`5k`<<;E(ido7M6lKTgWezNLq>U*=uz&s=cc$1%>VrAeOoUtA|T6gO4>UNqsdK=NF*8|~*sl&wI=x9-EGiq*aqV!(VVXA57 zw9*o6Ir8Lj1npUXvlevtn(_+^X5rzdR>#(}4YcB9O50q97%rW2me5_L=%ffYPUSRc z!vv?Kv>dH994Qi>U(a<0KF6NH5b16enCp+mw^Hb3Xs1^tThFpz!3QuN#}KBbww`(h z7GO)1olDqy6?T$()R7y%NYx*B0k_2IBiZ14&8|JPFxeMF{vW>HF-Vi3+ZOI=+qP}n zw(+!WcTd~4ZJX1!ZM&y!+uyt=&i!+~d(V%GjH;-NsEEv6nS1TERt|RHh!0>W4+4pp z1-*EzAM~i`+1f(VEHI8So`S`akPfPTfq*`l{Fz`hS%k#JS0cjT2mS0#QLGf=J?1`he3W*;m4)ce8*WFq1sdP=~$5RlH1EdWm|~dCvKOi4*I_96{^95p#B<(n!d?B z=o`0{t+&OMwKcxiBECznJcfH!fL(z3OvmxP#oWd48|mMjpE||zdiTBdWelj8&Qosv zZFp@&UgXuvJw5y=q6*28AtxZzo-UUpkRW%ne+Ylf!V-0+uQXBW=5S1o#6LXNtY5!I z%Rkz#(S8Pjz*P7bqB6L|M#Er{|QLae-Y{KA>`^} z@lPjeX>90X|34S-7}ZVXe{wEei1<{*e8T-Nbj8JmD4iwcE+Hg_zhkPVm#=@b$;)h6 z<<6y`nPa`f3I6`!28d@kdM{uJOgM%`EvlQ5B2bL)Sl=|y@YB3KeOzz=9cUW3clPAU z^sYc}xf9{4Oj?L5MOlYxR{+>w=vJjvbyO5}ptT(o6dR|ygO$)nVCvNGnq(6;bHlBd zl?w-|plD8spjDF03g5ip;W3Z z><0{BCq!Dw;h5~#1BuQilq*TwEu)qy50@+BE4bX28+7erX{BD4H)N+7U`AVEuREE8 z;X?~fyhF-x_sRfHIj~6f(+^@H)D=ngP;mwJjxhQUbUdzk8f94Ab%59-eRIq?ZKrwD z(BFI=)xrUlgu(b|hAysqK<}8bslmNNeD=#JW*}^~Nrswn^xw*nL@Tx!49bfJecV&KC2G4q5a!NSv)06A_5N3Y?veAz;Gv+@U3R% z)~UA8-0LvVE{}8LVDOHzp~2twReqf}ODIyXMM6=W>kL|OHcx9P%+aJGYi_Om)b!xe zF40Vntn0+VP>o<$AtP&JANjXBn7$}C@{+@3I@cqlwR2MdwGhVPxlTIcRVu@Ho-wO` z_~Or~IMG)A_`6-p)KPS@cT9mu9RGA>dVh5wY$NM9-^c@N=hcNaw4ITjm;iWSP^ZX| z)_XpaI61<+La+U&&%2a z0za$)-wZP@mwSELo#3!PGTt$uy0C(nTT@9NX*r3Ctw6J~7A(m#8fE)0RBd`TdKfAT zCf@$MAxjP`O(u9s@c0Fd@|}UQ6qp)O5Q5DPCeE6mSIh|Rj{$cAVIWsA=xPKVKxdhg zLzPZ`3CS+KIO;T}0Ip!fAUaNU>++ZJZRk@I(h<)RsJUhZ&Ru9*!4Ptn;gX^~4E8W^TSR&~3BAZc#HquXn)OW|TJ`CTahk+{qe`5+ixON^zA9IFd8)kc%*!AiLu z>`SFoZ5bW-%7}xZ>gpJcx_hpF$2l+533{gW{a7ce^B9sIdmLrI0)4yivZ^(Vh@-1q zFT!NQK$Iz^xu%|EOK=n>ug;(7J4OnS$;yWmq>A;hsD_0oAbLYhW^1Vdt9>;(JIYjf zdb+&f&D4@4AS?!*XpH>8egQvSVX`36jMd>$+RgI|pEg))^djhGSo&#lhS~9%NuWfX zDDH;3T*GzRT@5=7ibO>N-6_XPBYxno@mD_3I#rDD?iADxX`! zh*v8^i*JEMzyN#bGEBz7;UYXki*Xr(9xXax(_1qVW=Ml)kSuvK$coq2A(5ZGhs_pF z$*w}FbN6+QDseuB9=fdp_MTs)nQf!2SlROQ!gBJBCXD&@-VurqHj0wm@LWX-TDmS= z71M__vAok|@!qgi#H&H%Vg-((ZfxPAL8AI{x|VV!9)ZE}_l>iWk8UPTGHs*?u7RfP z5MC&=c6X;XlUzrz5q?(!eO@~* zoh2I*%J7dF!!_!vXoSIn5o|wj1#_>K*&CIn{qSaRc&iFVxt*^20ngCL;QonIS>I5^ zMw8HXm>W0PGd*}Ko)f|~dDd%;Wu_RWI_d;&2g6R3S63Uzjd7dn%Svu-OKpx*o|N>F zZg=-~qLb~VRLpv`k zWSdfHh@?dp=s_X`{yxOlxE$4iuyS;Z-x!*E6eqmEm*j2bE@=ZI0YZ5%Yj29!5+J$4h{s($nakA`xgbO8w zi=*r}PWz#lTL_DSAu1?f%-2OjD}NHXp4pXOsCW;DS@BC3h-q4_l`<))8WgzkdXg3! zs1WMt32kS2E#L0p_|x+x**TFV=gn`m9BWlzF{b%6j-odf4{7a4y4Uaef@YaeuPhU8 zHBvRqN^;$Jizy+ z=zW{E5<>2gp$pH{M@S*!sJVQU)b*J5*bX4h>5VJve#Q6ga}cQ&iL#=(u+KroWrxa%8&~p{WEUF0il=db;-$=A;&9M{Rq`ouZ5m%BHT6%st%saGsD6)fQgLN}x@d3q>FC;=f%O3Cyg=Ke@Gh`XW za@RajqOE9UB6eE=zhG%|dYS)IW)&y&Id2n7r)6p_)vlRP7NJL(x4UbhlcFXWT8?K=%s7;z?Vjts?y2+r|uk8Wt(DM*73^W%pAkZa1Jd zNoE)8FvQA>Z`eR5Z@Ig6kS5?0h;`Y&OL2D&xnnAUzQz{YSdh0k zB3exx%A2TyI)M*EM6htrxSlep!Kk(P(VP`$p0G~f$smld6W1r_Z+o?=IB@^weq>5VYsYZZR@` z&XJFxd5{|KPZmVOSxc@^%71C@;z}}WhbF9p!%yLj3j%YOlPL5s>7I3vj25 z@xmf=*z%Wb4;Va6SDk9cv|r*lhZ`(y_*M@>q;wrn)oQx%B(2A$9(74>;$zmQ!4fN; z>XurIk-7@wZys<+7XL@0Fhe-f%*=(weaQEdR9Eh6>Kl-EcI({qoZqyzziGwpg-GM#251sK_ z=3|kitS!j%;fpc@oWn65SEL73^N&t>Ix37xgs= zYG%eQDJc|rqHFia0!_sm7`@lvcv)gfy(+KXA@E{3t1DaZ$DijWAcA)E0@X?2ziJ{v z&KOYZ|DdkM{}t+@{@*6ge}m%xfjIxi%qh`=^2Rwz@w0cCvZ&Tc#UmCDbVwABrON^x zEBK43FO@weA8s7zggCOWhMvGGE`baZ62cC)VHyy!5Zbt%ieH+XN|OLbAFPZWyC6)p z4P3%8sq9HdS3=ih^0OOlqTPbKuzQ?lBEI{w^ReUO{V?@`ARsL|S*%yOS=Z%sF)>-y z(LAQdhgAcuF6LQjRYfdbD1g4o%tV4EiK&ElLB&^VZHbrV1K>tHTO{#XTo>)2UMm`2 z^t4s;vnMQgf-njU-RVBRw0P0-m#d-u`(kq7NL&2T)TjI_@iKuPAK-@oH(J8?%(e!0Ir$yG32@CGUPn5w4)+9@8c&pGx z+K3GKESI4*`tYlmMHt@br;jBWTei&(a=iYslc^c#RU3Q&sYp zSG){)V<(g7+8W!Wxeb5zJb4XE{I|&Y4UrFWr%LHkdQ;~XU zgy^dH-Z3lmY+0G~?DrC_S4@=>0oM8Isw%g(id10gWkoz2Q%7W$bFk@mIzTCcIB(K8 zc<5h&ZzCdT=9n-D>&a8vl+=ZF*`uTvQviG_bLde*k>{^)&0o*b05x$MO3gVLUx`xZ z43j+>!u?XV)Yp@MmG%Y`+COH2?nQcMrQ%k~6#O%PeD_WvFO~Kct za4XoCM_X!c5vhRkIdV=xUB3xI2NNStK*8_Zl!cFjOvp-AY=D;5{uXj}GV{LK1~IE2 z|KffUiBaStRr;10R~K2VVtf{TzM7FaPm;Y(zQjILn+tIPSrJh&EMf6evaBKIvi42-WYU9Vhj~3< zZSM-B;E`g_o8_XTM9IzEL=9Lb^SPhe(f(-`Yh=X6O7+6ALXnTcUFpI>ekl6v)ZQeNCg2 z^H|{SKXHU*%nBQ@I3It0m^h+6tvI@FS=MYS$ZpBaG7j#V@P2ZuYySbp@hA# ze(kc;P4i_-_UDP?%<6>%tTRih6VBgScKU^BV6Aoeg6Uh(W^#J^V$Xo^4#Ekp ztqQVK^g9gKMTHvV7nb64UU7p~!B?>Y0oFH5T7#BSW#YfSB@5PtE~#SCCg3p^o=NkMk$<8- z6PT*yIKGrvne7+y3}_!AC8NNeI?iTY(&nakN>>U-zT0wzZf-RuyZk^X9H-DT_*wk= z;&0}6LsGtfVa1q)CEUPlx#(ED@-?H<1_FrHU#z5^P3lEB|qsxEyn%FOpjx z3S?~gvoXy~L(Q{Jh6*i~=f%9kM1>RGjBzQh_SaIDfSU_9!<>*Pm>l)cJD@wlyxpBV z4Fmhc2q=R_wHCEK69<*wG%}mgD1=FHi4h!98B-*vMu4ZGW~%IrYSLGU{^TuseqVgV zLP<%wirIL`VLyJv9XG_p8w@Q4HzNt-o;U@Au{7%Ji;53!7V8Rv0^Lu^Vf*sL>R(;c zQG_ZuFl)Mh-xEIkGu}?_(HwkB2jS;HdPLSxVU&Jxy9*XRG~^HY(f0g8Q}iqnVmgjI zfd=``2&8GsycjR?M%(zMjn;tn9agcq;&rR!Hp z$B*gzHsQ~aXw8c|a(L^LW(|`yGc!qOnV(ZjU_Q-4z1&0;jG&vAKuNG=F|H?@m5^N@ zq{E!1n;)kNTJ>|Hb2ODt-7U~-MOIFo%9I)_@7fnX+eMMNh>)V$IXesJpBn|uo8f~#aOFytCT zf9&%MCLf8mp4kwHTcojWmM3LU=#|{3L>E}SKwOd?%{HogCZ_Z1BSA}P#O(%H$;z7XyJ^sjGX;j5 zrzp>|Ud;*&VAU3x#f{CKwY7Vc{%TKKqmB@oTHA9;>?!nvMA;8+Jh=cambHz#J18x~ zs!dF>$*AnsQ{{82r5Aw&^7eRCdvcgyxH?*DV5(I$qXh^zS>us*I66_MbL8y4d3ULj z{S(ipo+T3Ag!+5`NU2sc+@*m{_X|&p#O-SAqF&g_n7ObB82~$p%fXA5GLHMC+#qqL zdt`sJC&6C2)=juQ_!NeD>U8lDVpAOkW*khf7MCcs$A(wiIl#B9HM%~GtQ^}yBPjT@ z+E=|A!Z?A(rwzZ;T}o6pOVqHzTr*i;Wrc%&36kc@jXq~+w8kVrs;%=IFdACoLAcCAmhFNpbP8;s`zG|HC2Gv?I~w4ITy=g$`0qMQdkijLSOtX6xW%Z9Nw<;M- zMN`c7=$QxN00DiSjbVt9Mi6-pjv*j(_8PyV-il8Q-&TwBwH1gz1uoxs6~uU}PrgWB zIAE_I-a1EqlIaGQNbcp@iI8W1sm9fBBNOk(k&iLBe%MCo#?xI$%ZmGA?=)M9D=0t7 zc)Q0LnI)kCy{`jCGy9lYX%mUsDWwsY`;jE(;Us@gmWPqjmXL+Hu#^;k%eT>{nMtzj zsV`Iy6leTA8-PndszF;N^X@CJrTw5IIm!GPeu)H2#FQitR{1p;MasQVAG3*+=9FYK zw*k!HT(YQorfQj+1*mCV458(T5=fH`um$gS38hw(OqVMyunQ;rW5aPbF##A3fGH6h z@W)i9Uff?qz`YbK4c}JzQpuxuE3pcQO)%xBRZp{zJ^-*|oryTxJ-rR+MXJ)!f=+pp z10H|DdGd2exhi+hftcYbM0_}C0ZI-2vh+$fU1acsB-YXid7O|=9L!3e@$H*6?G*Zp z%qFB(sgl=FcC=E4CYGp4CN>=M8#5r!RU!u+FJVlH6=gI5xHVD&k;Ta*M28BsxfMV~ zLz+@6TxnfLhF@5=yQo^1&S}cmTN@m!7*c6z;}~*!hNBjuE>NLVl2EwN!F+)0$R1S! zR|lF%n!9fkZ@gPW|x|B={V6x3`=jS*$Pu0+5OWf?wnIy>Y1MbbGSncpKO0qE(qO=ts z!~@&!N`10S593pVQu4FzpOh!tvg}p%zCU(aV5=~K#bKi zHdJ1>tQSrhW%KOky;iW+O_n;`l9~omqM%sdxdLtI`TrJzN6BQz+7xOl*rM>xVI2~# z)7FJ^Dc{DC<%~VS?@WXzuOG$YPLC;>#vUJ^MmtbSL`_yXtNKa$Hk+l-c!aC7gn(Cg ze?YPYZ(2Jw{SF6MiO5(%_pTo7j@&DHNW`|lD`~{iH+_eSTS&OC*2WTT*a`?|9w1dh zh1nh@$a}T#WE5$7Od~NvSEU)T(W$p$s5fe^GpG+7fdJ9=enRT9$wEk+ZaB>G3$KQO zgq?-rZZnIv!p#>Ty~}c*Lb_jxJg$eGM*XwHUwuQ|o^}b3^T6Bxx{!?va8aC@-xK*H ztJBFvFfsSWu89%@b^l3-B~O!CXs)I6Y}y#0C0U0R0WG zybjroj$io0j}3%P7zADXOwHwafT#uu*zfM!oD$6aJx7+WL%t-@6^rD_a_M?S^>c;z zMK580bZXo1f*L$CuMeM4Mp!;P@}b~$cd(s5*q~FP+NHSq;nw3fbWyH)i2)-;gQl{S zZO!T}A}fC}vUdskGSq&{`oxt~0i?0xhr6I47_tBc`fqaSrMOzR4>0H^;A zF)hX1nfHs)%Zb-(YGX;=#2R6C{BG;k=?FfP?9{_uFLri~-~AJ;jw({4MU7e*d)?P@ zXX*GkNY9ItFjhwgAIWq7Y!ksbMzfqpG)IrqKx9q{zu%Mdl+{Dis#p9q`02pr1LG8R z@As?eG!>IoROgS!@J*to<27coFc1zpkh?w=)h9CbYe%^Q!Ui46Y*HO0mr% zEff-*$ndMNw}H2a5@BsGj5oFfd!T(F&0$<{GO!Qdd?McKkorh=5{EIjDTHU`So>8V zBA-fqVLb2;u7UhDV1xMI?y>fe3~4urv3%PX)lDw+HYa;HFkaLqi4c~VtCm&Ca+9C~ zge+67hp#R9`+Euq59WhHX&7~RlXn=--m8$iZ~~1C8cv^2(qO#X0?vl91gzUKBeR1J z^p4!!&7)3#@@X&2aF2-)1Ffcc^F8r|RtdL2X%HgN&XU-KH2SLCbpw?J5xJ*!F-ypZ zMG%AJ!Pr&}`LW?E!K~=(NJxuSVTRCGJ$2a*Ao=uUDSys!OFYu!Vs2IT;xQ6EubLIl z+?+nMGeQQhh~??0!s4iQ#gm3!BpMpnY?04kK375e((Uc7B3RMj;wE?BCoQGu=UlZt!EZ1Q*auI)dj3Jj{Ujgt zW5hd~-HWBLI_3HuO) zNrb^XzPsTIb=*a69wAAA3J6AAZZ1VsYbIG}a`=d6?PjM)3EPaDpW2YP$|GrBX{q*! z$KBHNif)OKMBCFP5>!1d=DK>8u+Upm-{hj5o|Wn$vh1&K!lVfDB&47lw$tJ?d5|=B z^(_9=(1T3Fte)z^>|3**n}mIX;mMN5v2F#l(q*CvU{Ga`@VMp#%rQkDBy7kYbmb-q z<5!4iuB#Q_lLZ8}h|hPODI^U6`gzLJre9u3k3c#%86IKI*^H-@I48Bi*@avYm4v!n0+v zWu{M{&F8#p9cx+gF0yTB_<2QUrjMPo9*7^-uP#~gGW~y3nfPAoV%amgr>PSyVAd@l)}8#X zR5zV6t*uKJZL}?NYvPVK6J0v4iVpwiN|>+t3aYiZSp;m0!(1`bHO}TEtWR1tY%BPB z(W!0DmXbZAsT$iC13p4f>u*ZAy@JoLAkJhzFf1#4;#1deO8#8d&89}en&z!W&A3++^1(;>0SB1*54d@y&9Pn;^IAf3GiXbfT`_>{R+Xv; zQvgL>+0#8-laO!j#-WB~(I>l0NCMt_;@Gp_f0#^c)t?&#Xh1-7RR0@zPyBz!U#0Av zT?}n({(p?p7!4S2ZBw)#KdCG)uPnZe+U|0{BW!m)9 zi_9$F?m<`2!`JNFv+w8MK_K)qJ^aO@7-Ig>cM4-r0bi=>?B_2mFNJ}aE3<+QCzRr*NA!QjHw# z`1OsvcoD0?%jq{*7b!l|L1+Tw0TTAM4XMq7*ntc-Ived>Sj_ZtS|uVdpfg1_I9knY z2{GM_j5sDC7(W&}#s{jqbybqJWyn?{PW*&cQIU|*v8YGOKKlGl@?c#TCnmnAkAzV- zmK={|1G90zz=YUvC}+fMqts0d4vgA%t6Jhjv?d;(Z}(Ep8fTZfHA9``fdUHkA+z3+ zhh{ohP%Bj?T~{i0sYCQ}uC#5BwN`skI7`|c%kqkyWIQ;!ysvA8H`b-t()n6>GJj6xlYDu~8qX{AFo$Cm3d|XFL=4uvc?Keb zzb0ZmMoXca6Mob>JqkNuoP>B2Z>D`Q(TvrG6m`j}-1rGP!g|qoL=$FVQYxJQjFn33lODt3Wb1j8VR zlR++vIT6^DtYxAv_hxupbLLN3e0%A%a+hWTKDV3!Fjr^cWJ{scsAdfhpI)`Bms^M6 zQG$waKgFr=c|p9Piug=fcJvZ1ThMnNhQvBAg-8~b1?6wL*WyqXhtj^g(Ke}mEfZVM zJuLNTUVh#WsE*a6uqiz`b#9ZYg3+2%=C(6AvZGc=u&<6??!slB1a9K)=VL zY9EL^mfyKnD zSJyYBc_>G;5RRnrNgzJz#Rkn3S1`mZgO`(r5;Hw6MveN(URf_XS-r58Cn80K)ArH4 z#Rrd~LG1W&@ttw85cjp8xV&>$b%nSXH_*W}7Ch2pg$$c0BdEo-HWRTZcxngIBJad> z;C>b{jIXjb_9Jis?NZJsdm^EG}e*pR&DAy0EaSGi3XWTa(>C%tz1n$u?5Fb z1qtl?;_yjYo)(gB^iQq?=jusF%kywm?CJP~zEHi0NbZ);$(H$w(Hy@{i>$wcVRD_X|w-~(0Z9BJyh zhNh;+eQ9BEIs;tPz%jSVnfCP!3L&9YtEP;svoj_bNzeGSQIAjd zBss@A;)R^WAu-37RQrM%{DfBNRx>v!G31Z}8-El9IOJlb_MSoMu2}GDYycNaf>uny z+8xykD-7ONCM!APry_Lw6-yT>5!tR}W;W`C)1>pxSs5o1z#j7%m=&=7O4hz+Lsqm` z*>{+xsabZPr&X=}G@obTb{nPTkccJX8w3CG7X+1+t{JcMabv~UNv+G?txRqXib~c^Mo}`q{$`;EBNJ;#F*{gvS12kV?AZ%O0SFB$^ zn+}!HbmEj}w{Vq(G)OGAzH}R~kS^;(-s&=ectz8vN!_)Yl$$U@HNTI-pV`LSj7Opu zTZ5zZ)-S_{GcEQPIQXLQ#oMS`HPu{`SQiAZ)m1at*Hy%3xma|>o`h%E%8BEbi9p0r zVjcsh<{NBKQ4eKlXU|}@XJ#@uQw*$4BxKn6#W~I4T<^f99~(=}a`&3(ur8R9t+|AQ zWkQx7l}wa48-jO@ft2h+7qn%SJtL%~890FG0s5g*kNbL3I&@brh&f6)TlM`K^(bhr zJWM6N6x3flOw$@|C@kPi7yP&SP?bzP-E|HSXQXG>7gk|R9BTj`e=4de9C6+H7H7n# z#GJeVs1mtHhLDmVO?LkYRQc`DVOJ_vdl8VUihO-j#t=0T3%Fc1f9F73ufJz*adn*p zc%&vi(4NqHu^R>sAT_0EDjVR8bc%wTz#$;%NU-kbDyL_dg0%TFafZwZ?5KZpcuaO54Z9hX zD$u>q!-9`U6-D`E#`W~fIfiIF5_m6{fvM)b1NG3xf4Auw;Go~Fu7cth#DlUn{@~yu z=B;RT*dp?bO}o%4x7k9v{r=Y@^YQ^UUm(Qmliw8brO^=NP+UOohLYiaEB3^DB56&V zK?4jV61B|1Uj_5fBKW;8LdwOFZKWp)g{B%7g1~DgO&N& z#lisxf?R~Z@?3E$Mms$$JK8oe@X`5m98V*aV6Ua}8Xs2#A!{x?IP|N(%nxsH?^c{& z@vY&R1QmQs83BW28qAmJfS7MYi=h(YK??@EhjL-t*5W!p z^gYX!Q6-vBqcv~ruw@oMaU&qp0Fb(dbVzm5xJN%0o_^@fWq$oa3X?9s%+b)x4w-q5Koe(@j6Ez7V@~NRFvd zfBH~)U5!ix3isg`6be__wBJp=1@yfsCMw1C@y+9WYD9_C%{Q~7^0AF2KFryfLlUP# zwrtJEcH)jm48!6tUcxiurAMaiD04C&tPe6DI0#aoqz#Bt0_7_*X*TsF7u*zv(iEfA z;$@?XVu~oX#1YXtceQL{dSneL&*nDug^OW$DSLF0M1Im|sSX8R26&)<0Fbh^*l6!5wfSu8MpMoh=2l z^^0Sr$UpZp*9oqa23fcCfm7`ya2<4wzJ`Axt7e4jJrRFVf?nY~2&tRL* zd;6_njcz01c>$IvN=?K}9ie%Z(BO@JG2J}fT#BJQ+f5LFSgup7i!xWRKw6)iITjZU z%l6hPZia>R!`aZjwCp}I zg)%20;}f+&@t;(%5;RHL>K_&7MH^S+7<|(SZH!u zznW|jz$uA`P9@ZWtJgv$EFp>)K&Gt+4C6#*khZQXS*S~6N%JDT$r`aJDs9|uXWdbg zBwho$phWx}x!qy8&}6y5Vr$G{yGSE*r$^r{}pw zVTZKvikRZ`J_IJrjc=X1uw?estdwm&bEahku&D04HD+0Bm~q#YGS6gp!KLf$A{%Qd z&&yX@Hp>~(wU{|(#U&Bf92+1i&Q*-S+=y=3pSZy$#8Uc$#7oiJUuO{cE6=tsPhwPe| zxQpK>`Dbka`V)$}e6_OXKLB%i76~4N*zA?X+PrhH<&)}prET;kel24kW%+9))G^JI zsq7L{P}^#QsZViX%KgxBvEugr>ZmFqe^oAg?{EI=&_O#e)F3V#rc z8$4}0Zr19qd3tE4#$3_f=Bbx9oV6VO!d3(R===i-7p=Vj`520w0D3W6lQfY48}!D* z&)lZMG;~er2qBoI2gsX+Ts-hnpS~NYRDtPd^FPzn!^&yxRy#CSz(b&E*tL|jIkq|l zf%>)7Dtu>jCf`-7R#*GhGn4FkYf;B$+9IxmqH|lf6$4irg{0ept__%)V*R_OK=T06 zyT_m-o@Kp6U{l5h>W1hGq*X#8*y@<;vsOFqEjTQXFEotR+{3}ODDnj;o0@!bB5x=N z394FojuGOtVKBlVRLtHp%EJv_G5q=AgF)SKyRN5=cGBjDWv4LDn$IL`*=~J7u&Dy5 zrMc83y+w^F&{?X(KOOAl-sWZDb{9X9#jrQtmrEXD?;h-}SYT7yM(X_6qksM=K_a;Z z3u0qT0TtaNvDER_8x*rxXw&C^|h{P1qxK|@pS7vdlZ#P z7PdB7MmC2}%sdzAxt>;WM1s0??`1983O4nFK|hVAbHcZ3x{PzytQLkCVk7hA!Lo` zEJH?4qw|}WH{dc4z%aB=0XqsFW?^p=X}4xnCJXK%c#ItOSjdSO`UXJyuc8bh^Cf}8 z@Ht|vXd^6{Fgai8*tmyRGmD_s_nv~r^Fy7j`Bu`6=G)5H$i7Q7lvQnmea&TGvJp9a|qOrUymZ$6G|Ly z#zOCg++$3iB$!6!>215A4!iryregKuUT344X)jQb3|9qY>c0LO{6Vby05n~VFzd?q zgGZv&FGlkiH*`fTurp>B8v&nSxNz)=5IF$=@rgND4d`!AaaX;_lK~)-U8la_Wa8i?NJC@BURO*sUW)E9oyv3RG^YGfN%BmxzjlT)bp*$<| zX3tt?EAy<&K+bhIuMs-g#=d1}N_?isY)6Ay$mDOKRh z4v1asEGWoAp=srraLW^h&_Uw|6O+r;wns=uwYm=JN4Q!quD8SQRSeEcGh|Eb5Jg8m zOT}u;N|x@aq)=&;wufCc^#)5U^VcZw;d_wwaoh9$p@Xrc{DD6GZUqZ ziC6OT^zSq@-lhbgR8B+e;7_Giv;DK5gn^$bs<6~SUadiosfewWDJu`XsBfOd1|p=q zE>m=zF}!lObA%ePey~gqU8S6h-^J2Y?>7)L2+%8kV}Gp=h`Xm_}rlm)SyUS=`=S7msKu zC|T!gPiI1rWGb1z$Md?0YJQ;%>uPLOXf1Z>N~`~JHJ!^@D5kSXQ4ugnFZ>^`zH8CAiZmp z6Ms|#2gcGsQ{{u7+Nb9sA?U>(0e$5V1|WVwY`Kn)rsnnZ4=1u=7u!4WexZD^IQ1Jk zfF#NLe>W$3m&C^ULjdw+5|)-BSHwpegdyt9NYC{3@QtMfd8GrIWDu`gd0nv-3LpGCh@wgBaG z176tikL!_NXM+Bv#7q^cyn9$XSeZR6#!B4JE@GVH zoobHZN_*RF#@_SVYKkQ_igme-Y5U}cV(hkR#k1c{bQNMji zU7aE`?dHyx=1`kOYZo_8U7?3-7vHOp`Qe%Z*i+FX!s?6huNp0iCEW-Z7E&jRWmUW_ z67j>)Ew!yq)hhG4o?^z}HWH-e=es#xJUhDRc4B51M4~E-l5VZ!&zQq`gWe`?}#b~7w1LH4Xa-UCT5LXkXQWheBa2YJYbyQ zl1pXR%b(KCXMO0OsXgl0P0Og<{(@&z1aokU-Pq`eQq*JYgt8xdFQ6S z6Z3IFSua8W&M#`~*L#r>Jfd6*BzJ?JFdBR#bDv$_0N!_5vnmo@!>vULcDm`MFU823 zpG9pqjqz^FE5zMDoGqhs5OMmC{Y3iVcl>F}5Rs24Y5B^mYQ;1T&ks@pIApHOdrzXF z-SdX}Hf{X;TaSxG_T$0~#RhqKISGKNK47}0*x&nRIPtmdwxc&QT3$8&!3fWu1eZ_P zJveQj^hJL#Sn!*4k`3}(d(aasl&7G0j0-*_2xtAnoX1@9+h zO#c>YQg60Z;o{Bi=3i7S`Ic+ZE>K{(u|#)9y}q*j8uKQ1^>+(BI}m%1v3$=4ojGBc zm+o1*!T&b}-lVvZqIUBc8V}QyFEgm#oyIuC{8WqUNV{Toz`oxhYpP!_p2oHHh5P@iB*NVo~2=GQm+8Yrkm2Xjc_VyHg1c0>+o~@>*Qzo zHVBJS>$$}$_4EniTI;b1WShX<5-p#TPB&!;lP!lBVBbLOOxh6FuYloD%m;n{r|;MU3!q4AVkua~fieeWu2 zQAQ$ue(IklX6+V;F1vCu-&V?I3d42FgWgsb_e^29ol}HYft?{SLf>DrmOp9o!t>I^ zY7fBCk+E8n_|apgM|-;^=#B?6RnFKlN`oR)`e$+;D=yO-(U^jV;rft^G_zl`n7qnM zL z*-Y4Phq+ZI1$j$F-f;`CD#|`-T~OM5Q>x}a>B~Gb3-+9i>Lfr|Ca6S^8g*{*?_5!x zH_N!SoRP=gX1?)q%>QTY!r77e2j9W(I!uAz{T`NdNmPBBUzi2{`XMB^zJGGwFWeA9 z{fk33#*9SO0)DjROug+(M)I-pKA!CX;IY(#gE!UxXVsa)X!UftIN98{pt#4MJHOhY zM$_l}-TJlxY?LS6Nuz1T<44m<4i^8k@D$zuCPrkmz@sdv+{ciyFJG2Zwy&%c7;atIeTdh!a(R^QXnu1Oq1b42*OQFWnyQ zWeQrdvP|w_idy53Wa<{QH^lFmEd+VlJkyiC>6B#s)F;w-{c;aKIm;Kp50HnA-o3lY z9B~F$gJ@yYE#g#X&3ADx&tO+P_@mnQTz9gv30_sTsaGXkfNYXY{$(>*PEN3QL>I!k zp)KibPhrfX3%Z$H6SY`rXGYS~143wZrG2;=FLj50+VM6soI~up_>fU(2Wl@{BRsMi zO%sL3x?2l1cXTF)k&moNsHfQrQ+wu(gBt{sk#CU=UhrvJIncy@tJX5klLjgMn>~h= zg|FR&;@eh|C7`>s_9c~0-{IAPV){l|Ts`i=)AW;d9&KPc3fMeoTS%8@V~D8*h;&(^>yjT84MM}=%#LS7shLAuuj(0VAYoozhWjq z4LEr?wUe2^WGwdTIgWBkDUJa>YP@5d9^Rs$kCXmMRxuF*YMVrn?0NFyPl}>`&dqZb z<5eqR=ZG3>n2{6v6BvJ`YBZeeTtB88TAY(x0a58EWyuf>+^|x8Qa6wA|1Nb_p|nA zWWa}|z8a)--Wj`LqyFk_a3gN2>5{Rl_wbW?#by7&i*^hRknK%jwIH6=dQ8*-_{*x0j^DUfMX0`|K@6C<|1cgZ~D(e5vBFFm;HTZF(!vT8=T$K+|F)x3kqzBV4-=p1V(lzi(s7jdu0>LD#N=$Lk#3HkG!a zIF<7>%B7sRNzJ66KrFV76J<2bdYhxll0y2^_rdG=I%AgW4~)1Nvz=$1UkE^J%BxLo z+lUci`UcU062os*=`-j4IfSQA{w@y|3}Vk?i;&SSdh8n+$iHA#%ERL{;EpXl6u&8@ zzg}?hkEOUOJt?ZL=pWZFJ19mI1@P=$U5*Im1e_8Z${JsM>Ov?nh8Z zP5QvI!{Jy@&BP48%P2{Jr_VgzW;P@7)M9n|lDT|Ep#}7C$&ud&6>C^5ZiwKIg2McPU(4jhM!BD@@L(Gd*Nu$ji(ljZ<{FIeW_1Mmf;76{LU z-ywN~=uNN)Xi6$<12A9y)K%X|(W0p|&>>4OXB?IiYr||WKDOJPxiSe01NSV-h24^L z_>m$;|C+q!Mj**-qQ$L-*++en(g|hw;M!^%_h-iDjFHLo-n3JpB;p?+o2;`*jpvJU zLY^lt)Un4joij^^)O(CKs@7E%*!w>!HA4Q?0}oBJ7Nr8NQ7QmY^4~jvf0-`%waOLn zdNjAPaC0_7c|RVhw)+71NWjRi!y>C+Bl;Z`NiL^zn2*0kmj5gyhCLCxts*cWCdRI| zjsd=sT5BVJc^$GxP~YF$-U{-?kW6r@^vHXB%{CqYzU@1>dzf#3SYedJG-Rm6^RB7s zGM5PR(yKPKR)>?~vpUIeTP7A1sc8-knnJk*9)3t^e%izbdm>Y=W{$wm(cy1RB-19i za#828DMBY+ps#7Y8^6t)=Ea@%Nkt)O6JCx|ybC;Ap}Z@Zw~*}3P>MZLPb4Enxz9Wf zssobT^(R@KuShj8>@!1M7tm|2%-pYYDxz-5`rCbaTCG5{;Uxm z*g=+H1X8{NUvFGzz~wXa%Eo};I;~`37*WrRU&K0dPSB$yk(Z*@K&+mFal^?c zurbqB-+|Kb5|sznT;?Pj!+kgFY1#Dr;_%A(GIQC{3ct|{*Bji%FNa6c-thbpBkA;U zURV!Dr&X{0J}iht#-Qp2=xzuh(fM>zRoiGrYl5ttw2#r34gC41CCOC31m~^UPTK@s z6;A@)7O7_%C)>bnAXerYuAHdE93>j2N}H${zEc6&SbZ|-fiG*-qtGuy-qDelH(|u$ zorf8_T6Zqe#Ub!+e3oSyrskt_HyW_^5lrWt#30l)tHk|j$@YyEkXUOV;6B51L;M@=NIWZXU;GrAa(LGxO%|im%7F<-6N;en0Cr zLH>l*y?pMwt`1*cH~LdBPFY_l;~`N!Clyfr;7w<^X;&(ZiVdF1S5e(+Q%60zgh)s4 zn2yj$+mE=miVERP(g8}G4<85^-5f@qxh2ec?n+$A_`?qN=iyT1?U@t?V6DM~BIlBB z>u~eXm-aE>R0sQy!-I4xtCNi!!qh?R1!kKf6BoH2GG{L4%PAz0{Sh6xpuyI%*~u)s z%rLuFl)uQUCBQAtMyN;%)zFMx4loh7uTfKeB2Xif`lN?2gq6NhWhfz0u5WP9J>=V2 zo{mLtSy&BA!mSzs&CrKWq^y40JF5a&GSXIi2= z{EYb59J4}VwikL4P=>+mc6{($FNE@e=VUwG+KV21;<@lrN`mnz5jYGASyvz7BOG_6(p^eTxD-4O#lROgon;R35=|nj#eHIfJBYPWG>H>`dHKCDZ3`R{-?HO0mE~(5_WYcFmp8sU?wr*UkAQiNDGc6T zA%}GOLXlOWqL?WwfHO8MB#8M8*~Y*gz;1rWWoVSXP&IbKxbQ8+s%4Jnt?kDsq7btI zCDr0PZ)b;B%!lu&CT#RJzm{l{2fq|BcY85`w~3LSK<><@(2EdzFLt9Y_`;WXL6x`0 zDoQ?=?I@Hbr;*VVll1Gmd8*%tiXggMK81a+T(5Gx6;eNb8=uYn z5BG-0g>pP21NPn>$ntBh>`*})Fl|38oC^9Qz>~MAazH%3Q~Qb!ALMf$srexgPZ2@&c~+hxRi1;}+)-06)!#Mq<6GhP z-Q?qmgo${aFBApb5p}$1OJKTClfi8%PpnczyVKkoHw7Ml9e7ikrF0d~UB}i3vizos zXW4DN$SiEV9{faLt5bHy2a>33K%7Td-n5C*N;f&ZqAg#2hIqEb(y<&f4u5BWJ>2^4 z414GosL=Aom#m&=x_v<0-fp1r%oVJ{T-(xnomNJ(Dryv zh?vj+%=II_nV+@NR+(!fZZVM&(W6{6%9cm+o+Z6}KqzLw{(>E86uA1`_K$HqINlb1 zKelh3-jr2I9V?ych`{hta9wQ2c9=MM`2cC{m6^MhlL2{DLv7C^j z$xXBCnDl_;l|bPGMX@*tV)B!c|4oZyftUlP*?$YU9C_eAsuVHJ58?)zpbr30P*C`T z7y#ao`uE-SOG(Pi+`$=e^mle~)pRrdwL5)N;o{gpW21of(QE#U6w%*C~`v-z0QqBML!!5EeYA5IQB0 z^l01c;L6E(iytN!LhL}wfwP7W9PNAkb+)Cst?qg#$n;z41O4&v+8-zPs+XNb-q zIeeBCh#ivnFLUCwfS;p{LC0O7tm+Sf9Jn)~b%uwP{%69;QC)Ok0t%*a5M+=;y8j=v z#!*pp$9@!x;UMIs4~hP#pnfVc!%-D<+wsG@R2+J&%73lK|2G!EQC)O05TCV=&3g)C!lT=czLpZ@Sa%TYuoE?v8T8`V;e$#Zf2_Nj6nvBgh1)2 GZ~q4|mN%#X literal 0 HcmV?d00001 diff --git a/bindings/jvm/gradle/wrapper/gradle-wrapper.properties b/bindings/jvm/gradle/wrapper/gradle-wrapper.properties new file mode 100644 index 0000000..df97d72 --- /dev/null +++ b/bindings/jvm/gradle/wrapper/gradle-wrapper.properties @@ -0,0 +1,7 @@ +distributionBase=GRADLE_USER_HOME +distributionPath=wrapper/dists +distributionUrl=https\://services.gradle.org/distributions/gradle-8.10.2-bin.zip +networkTimeout=10000 +validateDistributionUrl=true +zipStoreBase=GRADLE_USER_HOME +zipStorePath=wrapper/dists diff --git a/bindings/jvm/gradlew b/bindings/jvm/gradlew new file mode 100755 index 0000000..d95bf61 --- /dev/null +++ b/bindings/jvm/gradlew @@ -0,0 +1,252 @@ +#!/bin/sh + +# +# Copyright © 2015-2021 the original authors. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# https://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +# +# SPDX-License-Identifier: Apache-2.0 +# + +############################################################################## +# +# Gradle start up script for POSIX generated by Gradle. +# +# Important for running: +# +# (1) You need a POSIX-compliant shell to run this script. If your /bin/sh is +# noncompliant, but you have some other compliant shell such as ksh or +# bash, then to run this script, type that shell name before the whole +# command line, like: +# +# ksh Gradle +# +# Busybox and similar reduced shells will NOT work, because this script +# requires all of these POSIX shell features: +# * functions; +# * expansions «$var», «${var}», «${var:-default}», «${var+SET}», +# «${var#prefix}», «${var%suffix}», and «$( cmd )»; +# * compound commands having a testable exit status, especially «case»; +# * various built-in commands including «command», «set», and «ulimit». +# +# Important for patching: +# +# (2) This script targets any POSIX shell, so it avoids extensions provided +# by Bash, Ksh, etc; in particular arrays are avoided. +# +# The "traditional" practice of packing multiple parameters into a +# space-separated string is a well documented source of bugs and security +# problems, so this is (mostly) avoided, by progressively accumulating +# options in "$@", and eventually passing that to Java. +# +# Where the inherited environment variables (DEFAULT_JVM_OPTS, JAVA_OPTS, +# and GRADLE_OPTS) rely on word-splitting, this is performed explicitly; +# see the in-line comments for details. +# +# There are tweaks for specific operating systems such as AIX, CygWin, +# Darwin, MinGW, and NonStop. +# +# (3) This script is generated from the Groovy template +# https://github.com/gradle/gradle/blob/HEAD/platforms/jvm/plugins-application/src/main/resources/org/gradle/api/internal/plugins/unixStartScript.txt +# within the Gradle project. +# +# You can find Gradle at https://github.com/gradle/gradle/. +# +############################################################################## + +# Attempt to set APP_HOME + +# Resolve links: $0 may be a link +app_path=$0 + +# Need this for daisy-chained symlinks. +while + APP_HOME=${app_path%"${app_path##*/}"} # leaves a trailing /; empty if no leading path + [ -h "$app_path" ] +do + ls=$( ls -ld "$app_path" ) + link=${ls#*' -> '} + case $link in #( + /*) app_path=$link ;; #( + *) app_path=$APP_HOME$link ;; + esac +done + +# This is normally unused +# shellcheck disable=SC2034 +APP_BASE_NAME=${0##*/} +# Discard cd standard output in case $CDPATH is set (https://github.com/gradle/gradle/issues/25036) +APP_HOME=$( cd -P "${APP_HOME:-./}" > /dev/null && printf '%s +' "$PWD" ) || exit + +# Use the maximum available, or set MAX_FD != -1 to use that value. +MAX_FD=maximum + +warn () { + echo "$*" +} >&2 + +die () { + echo + echo "$*" + echo + exit 1 +} >&2 + +# OS specific support (must be 'true' or 'false'). +cygwin=false +msys=false +darwin=false +nonstop=false +case "$( uname )" in #( + CYGWIN* ) cygwin=true ;; #( + Darwin* ) darwin=true ;; #( + MSYS* | MINGW* ) msys=true ;; #( + NONSTOP* ) nonstop=true ;; +esac + +CLASSPATH=$APP_HOME/gradle/wrapper/gradle-wrapper.jar + + +# Determine the Java command to use to start the JVM. +if [ -n "$JAVA_HOME" ] ; then + if [ -x "$JAVA_HOME/jre/sh/java" ] ; then + # IBM's JDK on AIX uses strange locations for the executables + JAVACMD=$JAVA_HOME/jre/sh/java + else + JAVACMD=$JAVA_HOME/bin/java + fi + if [ ! -x "$JAVACMD" ] ; then + die "ERROR: JAVA_HOME is set to an invalid directory: $JAVA_HOME + +Please set the JAVA_HOME variable in your environment to match the +location of your Java installation." + fi +else + JAVACMD=java + if ! command -v java >/dev/null 2>&1 + then + die "ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. + +Please set the JAVA_HOME variable in your environment to match the +location of your Java installation." + fi +fi + +# Increase the maximum file descriptors if we can. +if ! "$cygwin" && ! "$darwin" && ! "$nonstop" ; then + case $MAX_FD in #( + max*) + # In POSIX sh, ulimit -H is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 + MAX_FD=$( ulimit -H -n ) || + warn "Could not query maximum file descriptor limit" + esac + case $MAX_FD in #( + '' | soft) :;; #( + *) + # In POSIX sh, ulimit -n is undefined. That's why the result is checked to see if it worked. + # shellcheck disable=SC2039,SC3045 + ulimit -n "$MAX_FD" || + warn "Could not set maximum file descriptor limit to $MAX_FD" + esac +fi + +# Collect all arguments for the java command, stacking in reverse order: +# * args from the command line +# * the main class name +# * -classpath +# * -D...appname settings +# * --module-path (only if needed) +# * DEFAULT_JVM_OPTS, JAVA_OPTS, and GRADLE_OPTS environment variables. + +# For Cygwin or MSYS, switch paths to Windows format before running java +if "$cygwin" || "$msys" ; then + APP_HOME=$( cygpath --path --mixed "$APP_HOME" ) + CLASSPATH=$( cygpath --path --mixed "$CLASSPATH" ) + + JAVACMD=$( cygpath --unix "$JAVACMD" ) + + # Now convert the arguments - kludge to limit ourselves to /bin/sh + for arg do + if + case $arg in #( + -*) false ;; # don't mess with options #( + /?*) t=${arg#/} t=/${t%%/*} # looks like a POSIX filepath + [ -e "$t" ] ;; #( + *) false ;; + esac + then + arg=$( cygpath --path --ignore --mixed "$arg" ) + fi + # Roll the args list around exactly as many times as the number of + # args, so each arg winds up back in the position where it started, but + # possibly modified. + # + # NB: a `for` loop captures its iteration list before it begins, so + # changing the positional parameters here affects neither the number of + # iterations, nor the values presented in `arg`. + shift # remove old arg + set -- "$@" "$arg" # push replacement arg + done +fi + + +# Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +DEFAULT_JVM_OPTS='-Dfile.encoding=UTF-8 "-Xmx64m" "-Xms64m"' + +# Collect all arguments for the java command: +# * DEFAULT_JVM_OPTS, JAVA_OPTS, JAVA_OPTS, and optsEnvironmentVar are not allowed to contain shell fragments, +# and any embedded shellness will be escaped. +# * For example: A user cannot expect ${Hostname} to be expanded, as it is an environment variable and will be +# treated as '${Hostname}' itself on the command line. + +set -- \ + "-Dorg.gradle.appname=$APP_BASE_NAME" \ + -classpath "$CLASSPATH" \ + org.gradle.wrapper.GradleWrapperMain \ + "$@" + +# Stop when "xargs" is not available. +if ! command -v xargs >/dev/null 2>&1 +then + die "xargs is not available" +fi + +# Use "xargs" to parse quoted args. +# +# With -n1 it outputs one arg per line, with the quotes and backslashes removed. +# +# In Bash we could simply go: +# +# readarray ARGS < <( xargs -n1 <<<"$var" ) && +# set -- "${ARGS[@]}" "$@" +# +# but POSIX shell has neither arrays nor command substitution, so instead we +# post-process each arg (as a line of input to sed) to backslash-escape any +# character that might be a shell metacharacter, then use eval to reverse +# that process (while maintaining the separation between arguments), and wrap +# the whole thing up as a single "set" statement. +# +# This will of course break if any of these variables contains a newline or +# an unmatched quote. +# + +eval "set -- $( + printf '%s\n' "$DEFAULT_JVM_OPTS $JAVA_OPTS $GRADLE_OPTS" | + xargs -n1 | + sed ' s~[^-[:alnum:]+,./:=@_]~\\&~g; ' | + tr '\n' ' ' + )" '"$@"' + +exec "$JAVACMD" "$@" diff --git a/bindings/jvm/gradlew.bat b/bindings/jvm/gradlew.bat new file mode 100644 index 0000000..640d686 --- /dev/null +++ b/bindings/jvm/gradlew.bat @@ -0,0 +1,94 @@ +@rem +@rem Copyright 2015 the original author or authors. +@rem +@rem Licensed under the Apache License, Version 2.0 (the "License"); +@rem you may not use this file except in compliance with the License. +@rem You may obtain a copy of the License at +@rem +@rem https://www.apache.org/licenses/LICENSE-2.0 +@rem +@rem Unless required by applicable law or agreed to in writing, software +@rem distributed under the License is distributed on an "AS IS" BASIS, +@rem WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +@rem See the License for the specific language governing permissions and +@rem limitations under the License. +@rem +@rem SPDX-License-Identifier: Apache-2.0 +@rem + +@if "%DEBUG%"=="" @echo off +@rem ########################################################################## +@rem +@rem Gradle startup script for Windows +@rem +@rem ########################################################################## + +@rem Set local scope for the variables with windows NT shell +if "%OS%"=="Windows_NT" setlocal + +set DIRNAME=%~dp0 +if "%DIRNAME%"=="" set DIRNAME=. +@rem This is normally unused +set APP_BASE_NAME=%~n0 +set APP_HOME=%DIRNAME% + +@rem Resolve any "." and ".." in APP_HOME to make it shorter. +for %%i in ("%APP_HOME%") do set APP_HOME=%%~fi + +@rem Add default JVM options here. You can also use JAVA_OPTS and GRADLE_OPTS to pass JVM options to this script. +set DEFAULT_JVM_OPTS=-Dfile.encoding=UTF-8 "-Xmx64m" "-Xms64m" + +@rem Find java.exe +if defined JAVA_HOME goto findJavaFromJavaHome + +set JAVA_EXE=java.exe +%JAVA_EXE% -version >NUL 2>&1 +if %ERRORLEVEL% equ 0 goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is not set and no 'java' command could be found in your PATH. 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:findJavaFromJavaHome +set JAVA_HOME=%JAVA_HOME:"=% +set JAVA_EXE=%JAVA_HOME%/bin/java.exe + +if exist "%JAVA_EXE%" goto execute + +echo. 1>&2 +echo ERROR: JAVA_HOME is set to an invalid directory: %JAVA_HOME% 1>&2 +echo. 1>&2 +echo Please set the JAVA_HOME variable in your environment to match the 1>&2 +echo location of your Java installation. 1>&2 + +goto fail + +:execute +@rem Setup the command line + +set CLASSPATH=%APP_HOME%\gradle\wrapper\gradle-wrapper.jar + + +@rem Execute Gradle +"%JAVA_EXE%" %DEFAULT_JVM_OPTS% %JAVA_OPTS% %GRADLE_OPTS% "-Dorg.gradle.appname=%APP_BASE_NAME%" -classpath "%CLASSPATH%" org.gradle.wrapper.GradleWrapperMain %* + +:end +@rem End local scope for the variables with windows NT shell +if %ERRORLEVEL% equ 0 goto mainEnd + +:fail +rem Set variable GRADLE_EXIT_CONSOLE if you need the _script_ return code instead of +rem the _cmd.exe /c_ return code! +set EXIT_CODE=%ERRORLEVEL% +if %EXIT_CODE% equ 0 set EXIT_CODE=1 +if not ""=="%GRADLE_EXIT_CONSOLE%" exit %EXIT_CODE% +exit /b %EXIT_CODE% + +:mainEnd +if "%OS%"=="Windows_NT" endlocal + +:omega diff --git a/bindings/jvm/settings.gradle.kts b/bindings/jvm/settings.gradle.kts new file mode 100644 index 0000000..2024db6 --- /dev/null +++ b/bindings/jvm/settings.gradle.kts @@ -0,0 +1 @@ +rootProject.name = "aibridge-jvm" diff --git a/bindings/jvm/src/main/java/io/aibridge/AibridgeException.java b/bindings/jvm/src/main/java/io/aibridge/AibridgeException.java new file mode 100644 index 0000000..4d40e9d --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/AibridgeException.java @@ -0,0 +1,88 @@ +package io.aibridge; + +/** + * AIBridge 异常基类(对应设计文档 9.3 节 JVM 异常映射)。 + * + *

FFI 错误模型:{@code aibridge_status_t} 返回码(0 成功 / 负数错误类别)+ + * {@code aibridge_last_error()} 线程局部 JSON(含 {@code code/message/details/retryable})。 + * 本类承载 last_error 的解析结果,子类按 {@code code} 分类。 + * + *

典型用法: + *

{@code
+ * try {
+ *     client.chat(req);
+ * } catch (RateLimitException e) {
+ *     if (e.isRetryable()) { /* 退避重试 *\/ }
+ * }
+ * }
+ */ +public class AibridgeException extends RuntimeException { + + /** 错误码(如 "rate_limit_error"、"validation_error"、"ffi_error") */ + private final String code; + /** 详细信息(JSON 值,可能为 null) */ + private final String details; + /** 是否可重试 */ + private final boolean retryable; + + public AibridgeException(String code, String message, String details, boolean retryable) { + super(message); + this.code = code; + this.details = details; + this.retryable = retryable; + } + + public AibridgeException(String code, String message, String details, boolean retryable, Throwable cause) { + super(message, cause); + this.code = code; + this.details = details; + this.retryable = retryable; + } + + /** 错误码(如 "rate_limit_error") */ + public String getCode() { + return code; + } + + /** 详细信息(原始 JSON 字符串,可能为 null 或 "null") */ + public String getDetails() { + return details; + } + + /** 是否可重试 */ + public boolean isRetryable() { + return retryable; + } + + @Override + public String toString() { + String retry = retryable ? " (retryable)" : ""; + return getClass().getSimpleName() + "[" + code + "]" + retry + ": " + getMessage(); + } + + // —— 错误码常量(与 aibridge-core error.rs code() 对齐)—— // + /** 认证错误 */ + public static final String CODE_AUTHENTICATION = "authentication_error"; + /** 限流错误 */ + public static final String CODE_RATE_LIMIT = "rate_limit_error"; + /** 参数校验错误 */ + public static final String CODE_VALIDATION = "validation_error"; + /** 模型不存在 */ + public static final String CODE_MODEL_NOT_FOUND = "model_not_found"; + /** API 调用错误(HTTP 4xx/5xx) */ + public static final String CODE_API = "api_error"; + /** 网络错误 */ + public static final String CODE_NETWORK = "network_error"; + /** 超时 */ + public static final String CODE_TIMEOUT = "timeout_error"; + /** 不支持的能力 */ + public static final String CODE_UNSUPPORTED_CAPABILITY = "unsupported_capability"; + /** Provider 不存在 */ + public static final String CODE_PROVIDER_NOT_FOUND = "provider_not_found"; + /** 音色不可用 */ + public static final String CODE_VOICE_NOT_AVAILABLE = "voice_not_available"; + /** 服务暂时不可用 */ + public static final String CODE_SERVICE_UNAVAILABLE = "service_unavailable"; + /** FFI 层通用错误(参数为空、JSON 解析失败、内部 panic 等) */ + public static final String CODE_FFI = "ffi_error"; +} diff --git a/bindings/jvm/src/main/java/io/aibridge/AibridgeNative.java b/bindings/jvm/src/main/java/io/aibridge/AibridgeNative.java new file mode 100644 index 0000000..2d10b13 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/AibridgeNative.java @@ -0,0 +1,179 @@ +package io.aibridge; + +import com.sun.jna.Library; +import com.sun.jna.Native; +import com.sun.jna.Pointer; +import com.sun.jna.Structure; +import com.sun.jna.ptr.PointerByReference; + +/** + * JNA Library 接口:声明所有 aibridge-ffi 的 C 函数。 + * + *

命名说明:本接口命名为 {@code AibridgeNative} 以避免与 JNA 的 + * {@link com.sun.jna.Native} 工具类同名冲突(后者提供 {@code Native.load})。 + * + *

JNA 通过 {@code Native.load} 加载 libaibridge(dylib/so/dll),并按 C ABI + * 调用。句柄(client/stream)用 {@link Pointer} 表示,复杂结构走 JSON 字符串边界, + * 二进制走 {@link AibridgeBytes}。 + * + *

库搜索路径: + *

    + *
  1. {@code jna.library.path} 系统属性
  2. + *
  3. {@code java.library.path} 系统属性
  4. + *
  5. {@code DYLD_LIBRARY_PATH} / {@code LD_LIBRARY_PATH} 环境变量
  6. + *
+ * + *

FFI 遗留约束(由 {@link Client} / {@link ChatStream} 负责): + *

    + *
  • {@link #aibridge_last_error()} 线程局部:调 FFI 同线程立即读,转存后抛异常
  • + *
  • {@link #aibridge_stream_next} 串行:同一 stream 不可并发 next
  • + *
  • {@link #aibridge_bytes_free} / {@link #aibridge_string_free} 必须调(RAII 封装)
  • + *
  • client/stream 句柄必须 destroy(RAII 封装)
  • + *
+ */ +public interface AibridgeNative extends Library { + + /** 库名(JNA 自动加平台前缀/后缀:libaibridge.dylib / libaibridge.so / aibridge.dll) */ + String LIBRARY_NAME = "aibridge"; + + /** 全局单例(JNA 内部线程安全) */ + AibridgeNative INSTANCE = Native.load(LIBRARY_NAME, AibridgeNative.class); + + // —— FFI 返回码常量(与 aibridge.h 的 AIBRIDGE_* 宏对齐)—— + int AIBRIDGE_OK = 0; + int AIBRIDGE_STREAM_CHUNK = 0; + int AIBRIDGE_STREAM_END = 1; + int AIBRIDGE_ERR_AUTHENTICATION = -1; + int AIBRIDGE_ERR_RATE_LIMIT = -2; + int AIBRIDGE_ERR_VALIDATION = -3; + int AIBRIDGE_ERR_MODEL_NOT_FOUND = -4; + int AIBRIDGE_ERR_API = -5; + int AIBRIDGE_ERR_NETWORK = -6; + int AIBRIDGE_ERR_TIMEOUT = -7; + int AIBRIDGE_ERR_UNSUPPORTED_CAPABILITY = -8; + int AIBRIDGE_ERR_PROVIDER_NOT_FOUND = -9; + int AIBRIDGE_ERR_VOICE_NOT_AVAILABLE = -10; + int AIBRIDGE_ERR_SERVICE_UNAVAILABLE = -11; + int AIBRIDGE_ERR_FFI = -100; + + /** + * 二进制缓冲结构(对应 C 的 {@code aibridge_bytes_t})。 + * + *

{@code ptr} 指向 Rust 分配的字节缓冲,{@code len} 为长度。 + * 调用方需通过 {@link #aibridge_bytes_free} 释放。 + * + *

JNA {@link Structure.ByReference} 用于 FFI 中 {@code aibridge_bytes_t**} + * 出参({@link #aibridge_client_speech} 的 {@code out_audio})。 + */ + @Structure.FieldOrder({"ptr", "len"}) + class AibridgeBytes extends Structure implements Structure.ByReference { + /** 指向字节数据的指针(Rust 分配) */ + public Pointer ptr; + /** 字节数据长度 */ + public long len; + + public AibridgeBytes() { + } + + /** 从结构体指针构造(用于解引用 {@code aibridge_bytes_t**} 出参) */ + public AibridgeBytes(Pointer p) { + super(p); + read(); + } + + /** 从结构体指针读取并拷贝为 Java 字节数组 */ + public byte[] toByteArray() { + if (ptr == null || len <= 0) { + return new byte[0]; + } + return ptr.getByteArray(0, (int) len); + } + } + + // —— 生命周期 —— // + + /** + * 创建客户端。 + * + * @param provider Provider 类型(如 "echo"、"openai"),UTF-8 C 字符串 + * @param configJson ClientOptions 的 JSON(可为 null,等价默认配置) + * @return client 指针;失败返回 null(错误写入 {@link #aibridge_last_error()}) + */ + Pointer aibridge_client_new(String provider, String configJson); + + /** + * 启动客户端(初始化适配器)。 + * + * @return 0 成功;负数为错误码 + */ + int aibridge_client_start(Pointer client); + + /** 释放客户端句柄(传 null 安全 no-op) */ + void aibridge_client_destroy(Pointer client); + + // —— 阻塞式调用 —— // + + /** + * 文本对话(阻塞)。 + * + * @param outResponseJson 写入 ChatCompletion 的 JSON(调用方需 {@link #aibridge_string_free}) + * @return 0 成功;负数错误码 + */ + int aibridge_client_chat(Pointer client, String requestJson, PointerByReference outResponseJson); + + /** + * 文字转语音(阻塞,二进制载荷走 {@link AibridgeBytes})。 + * + * @param outAudio 写入二进制音频缓冲(可为 null,调用方 {@link #aibridge_bytes_free}) + * @param outMetaJson 写入 SpeechResult(不含 audio_data)的 JSON({@link #aibridge_string_free}) + * @return 0 成功;负数错误码 + */ + int aibridge_client_speech( + Pointer client, + String requestJson, + PointerByReference outAudio, + PointerByReference outMetaJson); + + // —— 流式 —— // + + /** + * 创建流式对话(阻塞创建 stream 句柄)。 + * + * @param outStream 写入 stream 句柄 + * @return 0 成功;负数错误码 + */ + int aibridge_client_chat_stream( + Pointer client, + String requestJson, + PointerByReference outStream); + + /** + * 拉取下一个流式 chunk(阻塞,串行)。 + * + * @return 0=chunk({@code outChunkJson} 写入 JSON);1=EOF;负数=错误 + */ + int aibridge_stream_next(Pointer stream, PointerByReference outChunkJson); + + /** 释放 stream 句柄(传 null 安全 no-op,触发 Rust drop → tokio task abort) */ + void aibridge_stream_destroy(Pointer stream); + + // —— 错误查询 —— // + + /** + * 读取当前线程的 last_error(JSON 字符串)。 + * + *

返回指向线程局部缓冲的指针,调用方不应释放。仅在当前线程的下一次 + * FFI 调用前保证有效(thread_local 语义),故调用方需立即读取转存。 + * + * @return 错误 JSON 指针;当前线程无错误返回 null + */ + Pointer aibridge_last_error(); + + // —— 释放 —— // + + /** 释放 Rust 分配的 C 字符串(传 null 安全 no-op) */ + void aibridge_string_free(Pointer ptr); + + /** 释放 Rust 分配的二进制缓冲(传 null 安全 no-op) */ + void aibridge_bytes_free(AibridgeBytes ptr); +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ApiException.java b/bindings/jvm/src/main/java/io/aibridge/ApiException.java new file mode 100644 index 0000000..cc5bcbe --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ApiException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** API 调用错误(HTTP 4xx/5xx,对应 code = "api_error") */ +public class ApiException extends AibridgeException { + public ApiException(String message, String details, boolean retryable) { + super(CODE_API, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/AuthenticationException.java b/bindings/jvm/src/main/java/io/aibridge/AuthenticationException.java new file mode 100644 index 0000000..7545afe --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/AuthenticationException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 认证错误(API Key 无效/缺失等,对应 code = "authentication_error") */ +public class AuthenticationException extends AibridgeException { + public AuthenticationException(String message, String details, boolean retryable) { + super(CODE_AUTHENTICATION, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatChoice.java b/bindings/jvm/src/main/java/io/aibridge/ChatChoice.java new file mode 100644 index 0000000..54c6a36 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatChoice.java @@ -0,0 +1,22 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * 对话选项(对应 Rust {@code ChatChoice})。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChatChoice { + + /** 选项索引 */ + public int index; + /** 生成的回复消息 */ + public ChoiceMessage message; + /** 结束原因(stop / length / content_filter / tool_calls) */ + @JsonProperty("finish_reason") + public String finishReason; + + public ChatChoice() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatCompletion.java b/bindings/jvm/src/main/java/io/aibridge/ChatCompletion.java new file mode 100644 index 0000000..c12c79f --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatCompletion.java @@ -0,0 +1,38 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; + +/** + * 对话完成结果(对应 Rust {@code ChatCompletion})。 + * + *

由 {@code aibridge_client_chat} 的 {@code out_response_json} 反序列化得到。 + * {@link JsonIgnoreProperties} 忽略未知字段,保证前向兼容。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChatCompletion { + + /** 响应 ID */ + public String id; + /** 对象类型 */ + public String object; + /** 创建时间戳 */ + public long created; + /** 使用的模型 */ + public String model; + /** 回复选项列表 */ + public List choices; + /** Token 使用统计(可选) */ + public ChatUsage usage; + /** 服务层级(可选) */ + @JsonProperty("service_tier") + public String serviceTier; + /** 系统指纹(可选) */ + @JsonProperty("system_fingerprint") + public String systemFingerprint; + + public ChatCompletion() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatCompletionChunk.java b/bindings/jvm/src/main/java/io/aibridge/ChatCompletionChunk.java new file mode 100644 index 0000000..73e2703 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatCompletionChunk.java @@ -0,0 +1,34 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; + +/** + * 流式对话块(对应 Rust {@code ChatCompletionChunk})。 + * + *

由 {@code aibridge_stream_next} 的 {@code out_chunk_json} 反序列化得到。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChatCompletionChunk { + + public String id; + public String object; + public long created; + public String model; + public List choices; + public ChatUsage usage; + + public ChatCompletionChunk() { + } + + /** 取首个 delta 的增量内容(便捷方法,可能为 null) */ + public String firstDeltaContent() { + if (choices == null || choices.isEmpty()) { + return null; + } + DeltaMessage delta = choices.get(0).delta; + return delta == null ? null : delta.content; + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatCompletionDelta.java b/bindings/jvm/src/main/java/io/aibridge/ChatCompletionDelta.java new file mode 100644 index 0000000..c2d7915 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatCompletionDelta.java @@ -0,0 +1,19 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * 流式增量(对应 Rust {@code ChatCompletionDelta})。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChatCompletionDelta { + + public int index; + public DeltaMessage delta; + @JsonProperty("finish_reason") + public String finishReason; + + public ChatCompletionDelta() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatMessage.java b/bindings/jvm/src/main/java/io/aibridge/ChatMessage.java new file mode 100644 index 0000000..9e8cb72 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatMessage.java @@ -0,0 +1,50 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonCreator; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * 对话消息(对应 Rust {@code ChatMessage} tagged enum,{@code role} 作为 tag)。 + * + *

Rust 侧用 {@code #[serde(tag = "role", rename_all = "lowercase")]}, + * 故序列化形如 {@code {"role":"user","content":"..."}}。 + * + *

Java 侧用单一 POJO + {@code role} 字段简化(避免 tagged enum 的复杂映射)。 + * 当前覆盖 system/user/assistant 三种角色,足够 hello world 使用。 + */ +public class ChatMessage { + + /** 角色:system / user / assistant / tool */ + public String role; + /** 消息内容(user 的纯文本;assistant 的回复) */ + public String content; + /** 发送者名称(可选) */ + public String name; + + /** 默认构造(Jackson 反序列化需要) */ + public ChatMessage() { + } + + @JsonCreator + public ChatMessage( + @JsonProperty("role") String role, + @JsonProperty("content") String content) { + this.role = role; + this.content = content; + } + + /** 创建系统消息 */ + public static ChatMessage system(String content) { + return new ChatMessage("system", content); + } + + /** 创建用户消息(纯文本) */ + public static ChatMessage user(String content) { + return new ChatMessage("user", content); + } + + /** 创建助手消息 */ + public static ChatMessage assistant(String content) { + return new ChatMessage("assistant", content); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatRequest.java b/bindings/jvm/src/main/java/io/aibridge/ChatRequest.java new file mode 100644 index 0000000..bb65c86 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatRequest.java @@ -0,0 +1,73 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonInclude; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; + +/** + * 对话请求(对应 Rust {@code ChatRequest})。 + * + *

用 {@link Builder} 链式构造。仅声明 hello world 所需的核心字段, + * 其余可选参数(temperature/top_p/tools 等)按需扩展。 + * + *

序列化为 JSON 后通过 FFI 边界传给 {@code aibridge_client_chat} / + * {@code aibridge_client_chat_stream}。 + */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public class ChatRequest { + + /** 模型名称 */ + public String model; + /** 消息列表 */ + public List messages; + /** 温度系数 */ + public Double temperature; + /** 最大生成 token 数 */ + @JsonProperty("max_tokens") + public Integer maxTokens; + /** 是否流式输出(chat_stream 自动设置) */ + public Boolean stream; + + /** 默认构造(Jackson 反序列化需要) */ + public ChatRequest() { + } + + public ChatRequest(String model, List messages) { + this.model = model; + this.messages = messages; + } + + /** 创建 Builder */ + public static Builder builder(String model, List messages) { + return new Builder(model, messages); + } + + /** 链式构造器 */ + public static class Builder { + private final ChatRequest req; + + public Builder(String model, List messages) { + this.req = new ChatRequest(model, messages); + } + + public Builder temperature(double t) { + req.temperature = t; + return this; + } + + public Builder maxTokens(int n) { + req.maxTokens = n; + return this; + } + + public Builder stream(boolean s) { + req.stream = s; + return this; + } + + public ChatRequest build() { + return req; + } + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatStream.java b/bindings/jvm/src/main/java/io/aibridge/ChatStream.java new file mode 100644 index 0000000..d4194bb --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatStream.java @@ -0,0 +1,265 @@ +package io.aibridge; + +import com.fasterxml.jackson.databind.DeserializationFeature; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.sun.jna.Pointer; +import com.sun.jna.ptr.PointerByReference; + +import java.lang.ref.Cleaner; +import java.util.Iterator; +import java.util.NoSuchElementException; +import java.util.concurrent.Flow; +import java.util.concurrent.atomic.AtomicBoolean; + +/** + * 流式文本对话句柄(封装 native stream 句柄)。 + * + *

由 {@link Client#chatStream} 创建。提供两种消费方式: + *

    + *
  1. {@link Iterator}(阻塞遍历):{@code for (chunk : stream) { ... }}
  2. + *
  3. {@link Flow.Publisher}(反应式):{@link #subscribe}
  4. + *
+ * + *

FFI 遗留约束

+ *
    + *
  • stream_next 串行:同一 stream 不可并发 next。{@link #hasNext} / + * {@link #next} 用 synchronized 保证串行;反应式订阅在单线程串行拉取。
  • + *
  • chunk JSON 必须释放:每个 chunk 的 {@code out_chunk_json} 由 Rust 分配, + * 反序列化后立即 {@code aibridge_string_free}。
  • + *
  • stream 句柄必须 destroy:用 {@link Cleaner} 兜底,建议显式 close + * (try-with-resources)。destroy 触发 Rust drop → tokio task abort,实现取消语义。
  • + *
+ * + *

错误传递

+ *

流式过程中 {@code stream_next} 返回负数时,同线程立即读取 + * {@code aibridge_last_error()} 转存为异常。Iterator 模式下抛出;反应式模式下通过 + * {@link java.util.concurrent.Flow.Subscriber#onError} 传递。 + */ +public class ChatStream implements Iterator, Flow.Publisher, AutoCloseable { + + private static final Cleaner CLEANER = Cleaner.create(); + private static final ObjectMapper MAPPER = new ObjectMapper() + .configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); + + /** native stream 句柄(null 表示已结束/关闭) */ + private volatile Pointer handle; + /** 防止重复 close */ + private final AtomicBoolean closed = new AtomicBoolean(false); + /** Cleaner 兜底释放 */ + private final Cleaner.Cleanable cleanable; + + /** 预取的下一个 chunk(null 表示未预取或已结束) */ + private ChatCompletionChunk nextChunk; + /** 是否已到达 EOF */ + private boolean eof = false; + /** 流式错误(EOF 时若非 null 则表示以错误结束) */ + private AibridgeException error; + + public ChatStream(Pointer streamHandle) { + if (streamHandle == null) { + throw new AibridgeException(AibridgeException.CODE_FFI, + "stream 句柄为空", null, false); + } + this.handle = streamHandle; + this.cleanable = CLEANER.register(this, () -> AibridgeNative.INSTANCE.aibridge_stream_destroy(streamHandle)); + } + + // —— Iterator 模式(阻塞,串行)—— // + + @Override + public synchronized boolean hasNext() { + if (eof) { + return false; + } + if (nextChunk != null) { + return true; + } + // 预取下一个 chunk + pullNext(); + return nextChunk != null; + } + + @Override + public synchronized ChatCompletionChunk next() { + if (eof && nextChunk == null) { + if (error != null) { + throw error; + } + throw new NoSuchElementException("流已结束"); + } + if (nextChunk == null) { + pullNext(); + if (nextChunk == null) { + if (error != null) { + throw error; + } + throw new NoSuchElementException("流已结束"); + } + } + ChatCompletionChunk chunk = nextChunk; + nextChunk = null; + return chunk; + } + + /** 取流式错误(若以错误结束);正常结束返回 null。须在遍历结束后调用。 */ + public AibridgeException getError() { + return error; + } + + // —— Flow.Publisher 模式(反应式)—— // + + /** + * 反应式订阅:在单线程串行拉取 chunk 并推送给 subscriber。 + * + *

背压:subscriber 通过 {@code Subscription.request(n)} 申请 chunk,本实现用 + * {@link java.util.concurrent.Semaphore} 计数控制,每次 request 释放许可,拉取循环 + * 获取许可后才推下一个 chunk。取消({@code cancel})后停止拉取。 + * + *

FFI 串行约束:所有 {@code stream_next} 在单一拉取线程调用,天然串行。 + */ + @Override + public void subscribe(Flow.Subscriber subscriber) { + if (subscriber == null) { + throw new NullPointerException("subscriber 不能为空"); + } + final java.util.concurrent.Semaphore permits = new java.util.concurrent.Semaphore(0); + final java.util.concurrent.atomic.AtomicBoolean cancelled = new java.util.concurrent.atomic.AtomicBoolean(false); + + Flow.Subscription subscription = new Flow.Subscription() { + @Override + public void request(long n) { + if (n <= 0) { + subscriber.onError(new IllegalArgumentException("request 需正数: " + n)); + return; + } + permits.release((int) Math.min(n, Integer.MAX_VALUE)); + } + + @Override + public void cancel() { + cancelled.set(true); + permits.release(Integer.MAX_VALUE); + } + }; + + subscriber.onSubscribe(subscription); + + // 单线程串行拉取(保证 stream_next 串行,满足 FFI 约束) + Thread.ofVirtual().name("aibridge-stream-pull").start(() -> { + try { + while (!cancelled.get()) { + permits.acquire(); + if (cancelled.get()) { + break; + } + if (!hasNext()) { + // 流结束 + if (error != null) { + subscriber.onError(error); + } else { + subscriber.onComplete(); + } + return; + } + subscriber.onNext(next()); + } + } catch (InterruptedException e) { + Thread.currentThread().interrupt(); + subscriber.onError(e); + } catch (Exception e) { + subscriber.onError(e); + } + }); + } + + // —— 生命周期 —— // + + /** 关闭流,释放 stream 句柄。多次调用安全。 */ + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + return; + } + eof = true; + cleanable.clean(); + handle = null; + } + + // —— 内部:串行拉取下一个 chunk —— // + + /** + * 调用 {@code aibridge_stream_next} 拉取并解析下一个 chunk。 + * + *

同步块保证串行(FFI 遗留:同一 stream 不可并发 next)。 + * 结果写入 {@link #nextChunk};EOF 时设 {@link #eof};错误时设 {@link #error}。 + */ + private synchronized void pullNext() { + if (eof || handle == null) { + return; + } + PointerByReference outRef = new PointerByReference(); + int status = AibridgeNative.INSTANCE.aibridge_stream_next(handle, outRef); + + if (status == AibridgeNative.AIBRIDGE_STREAM_CHUNK) { + Pointer jsonPtr = outRef.getValue(); + if (jsonPtr == null) { + error = new AibridgeException(AibridgeException.CODE_FFI, + "stream_next 返回 chunk 但 out_chunk_json 为空", null, false); + eof = true; + return; + } + try { + String json = jsonPtr.getString(0, "UTF-8"); + nextChunk = parseChunk(json); + } finally { + // chunk JSON 由 Rust 分配,必须释放 + AibridgeNative.INSTANCE.aibridge_string_free(jsonPtr); + } + } else if (status == AibridgeNative.AIBRIDGE_STREAM_END) { + eof = true; + } else { + // 负数:错误,立即同线程读取 last_error 转存 + error = readLastError(); + eof = true; + } + } + + /** 解析 chunk JSON */ + private static ChatCompletionChunk parseChunk(String json) { + try { + return MAPPER.readValue(json, ChatCompletionChunk.class); + } catch (Exception e) { + throw new AibridgeException(AibridgeException.CODE_FFI, + "ChatCompletionChunk JSON 反序列化失败: " + e.getMessage() + + " (原始: " + json + ")", null, false, e); + } + } + + /** 读取 last_error 并映射为异常(同线程立即读取) */ + private static AibridgeException readLastError() { + Pointer errPtr = AibridgeNative.INSTANCE.aibridge_last_error(); + if (errPtr == null) { + return new AibridgeException(AibridgeException.CODE_FFI, + "未知错误(last_error 为空)", null, false); + } + String json = errPtr.getString(0, "UTF-8"); + try { + ErrorPayload payload = MAPPER.readValue(json, ErrorPayload.class); + String code = payload.code != null ? payload.code : AibridgeException.CODE_FFI; + String details = payload.details != null ? payload.details : "null"; + boolean retryable = Boolean.TRUE.equals(payload.retryable); + String message = payload.message != null ? payload.message : "(无错误消息)"; + return new AibridgeException(code, message, details, retryable); + } catch (Exception e) { + return new AibridgeException(AibridgeException.CODE_FFI, + "last_error JSON 解析失败: " + e.getMessage(), null, false, e); + } + } + + private static class ErrorPayload { + public String code; + public String message; + public String details; + public Boolean retryable; + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChatUsage.java b/bindings/jvm/src/main/java/io/aibridge/ChatUsage.java new file mode 100644 index 0000000..6b35a7e --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChatUsage.java @@ -0,0 +1,21 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * Token 使用统计(对应 Rust {@code ChatUsage})。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChatUsage { + + @JsonProperty("prompt_tokens") + public long promptTokens; + @JsonProperty("completion_tokens") + public long completionTokens; + @JsonProperty("total_tokens") + public long totalTokens; + + public ChatUsage() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ChoiceMessage.java b/bindings/jvm/src/main/java/io/aibridge/ChoiceMessage.java new file mode 100644 index 0000000..656b1c7 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ChoiceMessage.java @@ -0,0 +1,24 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; + +/** + * 完成结果中的消息(对应 Rust {@code ChoiceMessage})。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class ChoiceMessage { + + /** 角色(通常为 "assistant") */ + public String role; + /** 消息内容 */ + public String content; + /** 工具调用列表(可选) */ + @JsonProperty("tool_calls") + public List toolCalls; + + public ChoiceMessage() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/Client.java b/bindings/jvm/src/main/java/io/aibridge/Client.java new file mode 100644 index 0000000..51d91c4 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/Client.java @@ -0,0 +1,330 @@ +package io.aibridge; + +import com.fasterxml.jackson.databind.DeserializationFeature; +import com.fasterxml.jackson.databind.ObjectMapper; +import com.sun.jna.Pointer; +import com.sun.jna.ptr.PointerByReference; + +import java.lang.ref.Cleaner; +import java.util.List; +import java.util.concurrent.CompletableFuture; +import java.util.concurrent.Executor; +import java.util.concurrent.Executors; +import java.util.concurrent.atomic.AtomicBoolean; + +/** + * AIBridge 客户端(封装 native client 句柄)。 + * + *

通过 JNA 调用 aibridge-ffi 的 {@code aibridge_client_*} 函数,提供: + *

    + *
  • {@link #chat}:阻塞文本对话(异步版 {@link #chatAsync})
  • + *
  • {@link #chatStream}:流式文本对话(返回 {@link ChatStream})
  • + *
  • {@link #speech}:文字转语音(异步版 {@link #speechAsync})
  • + *
+ * + *

生命周期与内存管理

+ *

句柄用 {@link Cleaner} 兜底释放:即使忘记 {@link #close},GC 回收时也会调 + * {@code aibridge_client_destroy}。但建议显式 close 以尽早释放资源。 + * + *

错误处理(FFI 遗留:last_error 线程局部)

+ *

每个 FFI 调用失败后,在同一线程立即读取 {@code aibridge_last_error()} 转存 + * 为字符串,再映射为对应子类异常抛出。避免跨线程读取失效指针。 + * + *

异步

+ *

{@code *Async} 方法用 {@link CompletableFuture#supplyAsync} 在固定线程池上执行 + * 阻塞 FFI 调用(FFI 内部 {@code block_on} 会阻塞调用线程)。 + */ +public class Client implements AutoCloseable { + + /** 共享异步执行器(虚拟线程,适合阻塞 IO 密集的 FFI 调用) */ + private static final Executor ASYNC_EXECUTOR = + Executors.newThreadPerTaskExecutor(Thread.ofVirtual().name("aibridge-ffi-", 0).factory()); + + private static final Cleaner CLEANER = Cleaner.create(); + private static final ObjectMapper MAPPER = new ObjectMapper() + .configure(DeserializationFeature.FAIL_ON_UNKNOWN_PROPERTIES, false); + + /** native client 句柄(null 表示已关闭) */ + private volatile Pointer handle; + /** 防止重复 close */ + private final AtomicBoolean closed = new AtomicBoolean(false); + /** Cleaner 注册的清理动作(兜底释放句柄) */ + private final Cleaner.Cleanable cleanable; + + /** + * 创建客户端。 + * + * @param provider Provider 类型(如 "echo"、"openai") + * @param configJson ClientOptions 的 JSON(可为 null,等价默认配置) + * @throws AibridgeException 创建失败(如 provider 不存在、config 非法) + */ + public Client(String provider, String configJson) { + Pointer ptr = AibridgeNative.INSTANCE.aibridge_client_new(provider, configJson); + if (ptr == null) { + // client_new 失败时 last_error 已写入,立即同线程读取转存 + throw readLastError(); + } + this.handle = ptr; + // Cleaner 兜底:GC 回收时调 destroy(传 null 安全 no-op) + this.cleanable = CLEANER.register(this, () -> AibridgeNative.INSTANCE.aibridge_client_destroy(ptr)); + } + + /** 创建客户端(默认配置) */ + public Client(String provider) { + this(provider, null); + } + + /** + * 启动客户端(初始化适配器)。 + * + * @throws AibridgeException 启动失败 + */ + public void start() { + Pointer ptr = requireHandle(); + int status = AibridgeNative.INSTANCE.aibridge_client_start(ptr); + if (status != AibridgeNative.AIBRIDGE_OK) { + throw readLastError(); + } + } + + /** + * 文本对话(阻塞)。 + * + * @param request 对话请求 + * @return 对话完成结果 + * @throws AibridgeException 调用失败 + */ + public ChatCompletion chat(ChatRequest request) { + Pointer ptr = requireHandle(); + String requestJson = writeJson(request); + + PointerByReference outRef = new PointerByReference(); + int status = AibridgeNative.INSTANCE.aibridge_client_chat(ptr, requestJson, outRef); + if (status != AibridgeNative.AIBRIDGE_OK) { + throw readLastError(); + } + // out_response_json 由 Rust 分配,必须 string_free + Pointer jsonPtr = outRef.getValue(); + try { + String json = jsonPtr.getString(0, "UTF-8"); + return parseJson(json, ChatCompletion.class); + } finally { + AibridgeNative.INSTANCE.aibridge_string_free(jsonPtr); + } + } + + /** + * 文本对话(异步)。 + * + * @return CompletableFuture,正常完成时携带 ChatCompletion + */ + public CompletableFuture chatAsync(ChatRequest request) { + return CompletableFuture.supplyAsync(() -> chat(request), ASYNC_EXECUTOR); + } + + /** + * 流式文本对话(创建 stream 句柄并启动后台拉取)。 + * + *

返回的 {@link ChatStream} 实现了 {@link java.util.Iterator},可阻塞遍历 chunk。 + * 使用完毕务必 {@link ChatStream#close()} 释放 stream 句柄(try-with-resources 推荐)。 + * + * @param request 对话请求({@code stream} 字段会被强制设为 true) + * @return 流式迭代器 + * @throws AibridgeException 创建流失败 + */ + public ChatStream chatStream(ChatRequest request) { + Pointer ptr = requireHandle(); + // 强制 stream=true(语义清晰,避免调用方遗漏) + if (request.stream == null) { + request.stream = true; + } + String requestJson = writeJson(request); + + PointerByReference outRef = new PointerByReference(); + int status = AibridgeNative.INSTANCE.aibridge_client_chat_stream(ptr, requestJson, outRef); + if (status != AibridgeNative.AIBRIDGE_OK) { + throw readLastError(); + } + Pointer streamPtr = outRef.getValue(); + if (streamPtr == null) { + throw new AibridgeException(AibridgeException.CODE_FFI, + "chat_stream 返回成功但 stream 句柄为空", null, false); + } + return new ChatStream(streamPtr); + } + + /** + * 文字转语音(阻塞)。 + * + * @param request 语音请求 + * @return 完整结果(meta + 二进制音频) + * @throws AibridgeException 调用失败 + */ + public SpeechResultFull speech(SpeechRequest request) { + Pointer ptr = requireHandle(); + String requestJson = writeJson(request); + + PointerByReference outAudioRef = new PointerByReference(); + PointerByReference outMetaRef = new PointerByReference(); + int status = AibridgeNative.INSTANCE.aibridge_client_speech(ptr, requestJson, outAudioRef, outMetaRef); + if (status != AibridgeNative.AIBRIDGE_OK) { + throw readLastError(); + } + + // 二进制音频(可为 null:Provider 仅返回 base64/url 时) + Pointer audioPtr = outAudioRef.getValue(); + byte[] audioData; + if (audioPtr == null) { + audioData = new byte[0]; + } else { + AibridgeNative.AibridgeBytes bytes = new AibridgeNative.AibridgeBytes(audioPtr); + bytes.read(); + audioData = bytes.toByteArray(); + AibridgeNative.INSTANCE.aibridge_bytes_free(bytes); + } + + // meta JSON(SpeechResult,audio_data 被 skip) + Pointer metaPtr = outMetaRef.getValue(); + try { + String json = metaPtr.getString(0, "UTF-8"); + SpeechResult meta = parseJson(json, SpeechResult.class); + return new SpeechResultFull(meta, audioData); + } finally { + AibridgeNative.INSTANCE.aibridge_string_free(metaPtr); + } + } + + /** + * 文字转语音(异步)。 + * + * @return CompletableFuture,正常完成时携带 SpeechResultFull + */ + public CompletableFuture speechAsync(SpeechRequest request) { + return CompletableFuture.supplyAsync(() -> speech(request), ASYNC_EXECUTOR); + } + + /** 关闭客户端,释放 native 句柄。多次调用安全。 */ + @Override + public void close() { + if (!closed.compareAndSet(false, true)) { + return; + } + // 取消 Cleaner 兜底(避免重复 destroy,destroy(null) 亦安全) + cleanable.clean(); + handle = null; + } + + // —— 内部辅助 —— // + + /** 校验句柄有效,否则抛 ffi_error */ + private Pointer requireHandle() { + Pointer ptr = handle; + if (ptr == null) { + throw new AibridgeException(AibridgeException.CODE_FFI, + "client 句柄为空(已 close 或未初始化)", null, false); + } + return ptr; + } + + /** + * 读取当前线程的 last_error 并映射为对应子类异常。 + * + *

FFI 遗留:last_error 是线程局部,必须在与触发错误的 FFI 调用相同的线程立即读取。 + * 本方法在 FFI 调用失败后立即被调用,故线程一致。 + */ + private static AibridgeException readLastError() { + Pointer errPtr = AibridgeNative.INSTANCE.aibridge_last_error(); + if (errPtr == null) { + return new AibridgeException(AibridgeException.CODE_FFI, + "未知错误(last_error 为空)", null, false); + } + // 立即转存(指针仅在下次 FFI 调用前有效) + String json = errPtr.getString(0, "UTF-8"); + return parseError(json); + } + + /** 解析 last_error JSON 并映射为对应子类异常 */ + private static AibridgeException parseError(String json) { + try { + ErrorPayload payload = MAPPER.readValue(json, ErrorPayload.class); + String code = payload.code != null ? payload.code : AibridgeException.CODE_FFI; + String details = payload.details != null ? payload.details : "null"; + boolean retryable = Boolean.TRUE.equals(payload.retryable); + String message = payload.message != null ? payload.message : "(无错误消息)"; + return mapToException(code, message, details, retryable); + } catch (Exception e) { + // JSON 解析失败,回退 ffi_error + return new AibridgeException(AibridgeException.CODE_FFI, + "last_error JSON 解析失败: " + e.getMessage() + " (原始: " + json + ")", + null, false, e); + } + } + + /** 按 code 映射到具体子类(与 aibridge-core error.rs code() 对齐) */ + private static AibridgeException mapToException(String code, String message, String details, boolean retryable) { + switch (code) { + case AibridgeException.CODE_AUTHENTICATION: + return new AuthenticationException(message, details, retryable); + case AibridgeException.CODE_RATE_LIMIT: + return new RateLimitException(message, details, retryable); + case AibridgeException.CODE_VALIDATION: + return new ValidationException(message, details, retryable); + case AibridgeException.CODE_MODEL_NOT_FOUND: + return new ModelNotFoundException(message, details, retryable); + case AibridgeException.CODE_API: + return new ApiException(message, details, retryable); + case AibridgeException.CODE_NETWORK: + return new NetworkException(message, details, retryable); + case AibridgeException.CODE_TIMEOUT: + return new TimeoutException(message, details, retryable); + case AibridgeException.CODE_UNSUPPORTED_CAPABILITY: + return new UnsupportedCapabilityException(message, details, retryable); + case AibridgeException.CODE_PROVIDER_NOT_FOUND: + return new ProviderNotFoundException(message, details, retryable); + case AibridgeException.CODE_VOICE_NOT_AVAILABLE: + return new VoiceNotAvailableException(message, details, retryable); + case AibridgeException.CODE_SERVICE_UNAVAILABLE: + return new ServiceUnavailableException(message, details, retryable); + default: + return new AibridgeException(code, message, details, retryable); + } + } + + /** 序列化为 JSON 字符串 */ + private static String writeJson(Object obj) { + try { + return MAPPER.writeValueAsString(obj); + } catch (Exception e) { + throw new AibridgeException(AibridgeException.CODE_FFI, + "JSON 序列化失败: " + e.getMessage(), null, false, e); + } + } + + /** 反序列化 JSON */ + private static T parseJson(String json, Class type) { + try { + return MAPPER.readValue(json, type); + } catch (Exception e) { + throw new AibridgeException(AibridgeException.CODE_FFI, + type.getSimpleName() + " JSON 反序列化失败: " + e.getMessage() + + " (原始: " + json + ")", null, false, e); + } + } + + /** last_error JSON 载荷(内部解析用) */ + private static class ErrorPayload { + public String code; + public String message; + public String details; + public Boolean retryable; + } + + /** 便捷构造:user 消息单条 */ + public static List userMessages(String... texts) { + java.util.ArrayList list = new java.util.ArrayList<>(); + for (String t : texts) { + list.add(ChatMessage.user(t)); + } + return list; + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/DeltaMessage.java b/bindings/jvm/src/main/java/io/aibridge/DeltaMessage.java new file mode 100644 index 0000000..392dc2b --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/DeltaMessage.java @@ -0,0 +1,24 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.List; + +/** + * 流式增量消息(对应 Rust {@code DeltaMessage})。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class DeltaMessage { + + /** 角色(首个块通常为 "assistant") */ + public String role; + /** 增量内容 */ + public String content; + /** 工具调用增量(可选) */ + @JsonProperty("tool_calls") + public List toolCalls; + + public DeltaMessage() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/Hello.java b/bindings/jvm/src/main/java/io/aibridge/Hello.java new file mode 100644 index 0000000..b0b2ffc --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/Hello.java @@ -0,0 +1,133 @@ +package io.aibridge; + +import java.util.List; + +/** + * AIBridge JVM 绑定 hello world(阶段 0.6 管线验证)。 + * + *

使用 echo 适配器(免认证)端到端验证: + *

    + *
  1. {@code chat}:echo-chat 模型,回显最后一条 user 消息 + " [echo]"(期望 "hello [echo]")
  2. + *
  3. {@code chatStream}:echo-chat 模型,3 个 chunk(role / 前半段 / 后半段+finish)
  4. + *
  5. {@code speech}:echo-tts 模型,返固定 15 字节音频
  6. + *
+ * + *

运行: + *

{@code
+ * ./gradlew run
+ * # 或:./gradlew build && java -Djava.library.path=../../target/debug -jar build/libs/aibridge-jvm-*.jar
+ * }
+ * 库搜索路径默认指向 {@code ../../target/debug}(见 build.gradle.kts)。 + */ +public class Hello { + + public static void main(String[] args) { + System.out.println("=== AIBridge JVM 绑定 Hello World ==="); + System.out.println("JNA 库路径: " + System.getProperty("jna.library.path", "(默认)")); + System.out.println(); + + // try-with-resources 保证 close 释放 client 句柄 + try (Client client = new Client("echo")) { + client.start(); + System.out.println("[OK] client 创建并启动成功"); + + // 1. 阻塞 chat + testChat(client); + + // 2. 流式 chatStream + testChatStream(client); + + // 3. 文字转语音 speech + testSpeech(client); + + System.out.println(); + System.out.println("=== 全部测试通过 ==="); + } catch (AibridgeException e) { + System.err.println("[FAIL] AIBridge 错误: " + e); + System.err.println(" code=" + e.getCode() + " retryable=" + e.isRetryable()); + System.err.println(" details=" + e.getDetails()); + System.exit(1); + } catch (Exception e) { + System.err.println("[FAIL] 意外错误: " + e); + e.printStackTrace(); + System.exit(1); + } + } + + /** 测试阻塞 chat:期望回显 "hello [echo]" */ + private static void testChat(Client client) { + System.out.println("--- 测试 1:阻塞 chat ---"); + ChatRequest req = ChatRequest.builder( + "echo-chat", + List.of(ChatMessage.user("hello"))) + .build(); + ChatCompletion resp = client.chat(req); + String content = resp.choices.get(0).message.content; + System.out.println(" chat 响应 id=" + resp.id); + System.out.println(" choices[0].message.content = \"" + content + "\""); + if (!"hello [echo]".equals(content)) { + throw new IllegalStateException("chat 回显不符:期望 \"hello [echo]\",实际 \"" + content + "\""); + } + System.out.println("[OK] chat 回显正确"); + System.out.println(); + } + + /** 测试流式 chatStream:期望 3 个 chunk */ + private static void testChatStream(Client client) { + System.out.println("--- 测试 2:流式 chatStream ---"); + ChatRequest req = ChatRequest.builder( + "echo-chat", + List.of(ChatMessage.user("hello"))) + .stream(true) + .build(); + + int chunkCount = 0; + StringBuilder assembled = new StringBuilder(); + // try-with-resources 保证 close 释放 stream 句柄 + try (ChatStream stream = client.chatStream(req)) { + while (stream.hasNext()) { + ChatCompletionChunk chunk = stream.next(); + chunkCount++; + String deltaContent = chunk.firstDeltaContent(); + String role = chunk.choices.isEmpty() ? null : chunk.choices.get(0).delta.role; + String finish = chunk.choices.isEmpty() ? null : chunk.choices.get(0).finishReason; + System.out.println(" chunk[" + chunkCount + "] role=" + role + + " content=" + (deltaContent == null ? "(null)" : "\"" + deltaContent + "\"") + + " finish=" + finish); + if (deltaContent != null) { + assembled.append(deltaContent); + } + } + if (stream.getError() != null) { + throw stream.getError(); + } + } + + System.out.println(" 流式拼接内容 = \"" + assembled + "\""); + System.out.println(" chunk 总数 = " + chunkCount); + if (chunkCount != 3) { + throw new IllegalStateException("chunk 数不符:期望 3,实际 " + chunkCount); + } + if (!"hello [echo]".equals(assembled.toString())) { + throw new IllegalStateException("流式拼接不符:期望 \"hello [echo]\",实际 \"" + assembled + "\""); + } + System.out.println("[OK] chatStream 3 个 chunk + 拼接正确"); + System.out.println(); + } + + /** 测试 speech:期望 15 字节音频 */ + private static void testSpeech(Client client) { + System.out.println("--- 测试 3:文字转语音 speech ---"); + SpeechRequest req = SpeechRequest.builder("echo-tts", "hello", "alloy").build(); + SpeechResultFull result = client.speech(req); + int len = result.audioLength(); + System.out.println(" audioData.length = " + len); + System.out.println(" content_type = " + result.getMeta().contentType); + System.out.println(" format = " + result.getMeta().format); + if (len != 15) { + throw new IllegalStateException("音频长度不符:期望 15,实际 " + len); + } + System.out.println("[OK] speech 返回 15 字节音频"); + System.out.println(); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ModelNotFoundException.java b/bindings/jvm/src/main/java/io/aibridge/ModelNotFoundException.java new file mode 100644 index 0000000..8018182 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ModelNotFoundException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 模型不存在(对应 code = "model_not_found") */ +public class ModelNotFoundException extends AibridgeException { + public ModelNotFoundException(String message, String details, boolean retryable) { + super(CODE_MODEL_NOT_FOUND, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/NetworkException.java b/bindings/jvm/src/main/java/io/aibridge/NetworkException.java new file mode 100644 index 0000000..8496d6c --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/NetworkException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 网络错误(对应 code = "network_error") */ +public class NetworkException extends AibridgeException { + public NetworkException(String message, String details, boolean retryable) { + super(CODE_NETWORK, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ProviderNotFoundException.java b/bindings/jvm/src/main/java/io/aibridge/ProviderNotFoundException.java new file mode 100644 index 0000000..005a5c5 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ProviderNotFoundException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** Provider 不存在(对应 code = "provider_not_found") */ +public class ProviderNotFoundException extends AibridgeException { + public ProviderNotFoundException(String message, String details, boolean retryable) { + super(CODE_PROVIDER_NOT_FOUND, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/RateLimitException.java b/bindings/jvm/src/main/java/io/aibridge/RateLimitException.java new file mode 100644 index 0000000..fabaebb --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/RateLimitException.java @@ -0,0 +1,12 @@ +package io.aibridge; + +/** + * 限流错误(对应 code = "rate_limit_error")。 + * + *

{@code details} 可能包含 {@code retry_after}(建议等待秒数)。 + */ +public class RateLimitException extends AibridgeException { + public RateLimitException(String message, String details, boolean retryable) { + super(CODE_RATE_LIMIT, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ServiceUnavailableException.java b/bindings/jvm/src/main/java/io/aibridge/ServiceUnavailableException.java new file mode 100644 index 0000000..c8c739f --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ServiceUnavailableException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 服务暂时不可用(对应 code = "service_unavailable") */ +public class ServiceUnavailableException extends AibridgeException { + public ServiceUnavailableException(String message, String details, boolean retryable) { + super(CODE_SERVICE_UNAVAILABLE, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/SpeechRequest.java b/bindings/jvm/src/main/java/io/aibridge/SpeechRequest.java new file mode 100644 index 0000000..898a9a8 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/SpeechRequest.java @@ -0,0 +1,63 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonInclude; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * 文字转语音请求(对应 Rust {@code SpeechRequest})。 + * + *

注意 {@code voice} 在 Rust 侧是 {@code VoiceSpec}(含 {@code voices} 数组), + * 支持候选列表自动降级。这里用 {@link VoiceSpec} 嵌套表示。 + */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public class SpeechRequest { + + /** 模型名称(如 "tts-1"、"echo-tts") */ + public String model; + /** 要合成的文本 */ + public String input; + /** 音色规格(候选列表) */ + public VoiceSpec voice; + /** 音频输出格式("mp3" / "opus" / "wav" 等) */ + @JsonProperty("response_format") + public String responseFormat; + /** 语速(0.25-4.0) */ + public Double speed; + + public SpeechRequest() { + } + + public SpeechRequest(String model, String input, String voice) { + this.model = model; + this.input = input; + this.voice = VoiceSpec.single(voice); + } + + /** 创建 Builder(单音色) */ + public static Builder builder(String model, String input, String voice) { + return new Builder(model, input, voice); + } + + /** 链式构造器 */ + public static class Builder { + private final SpeechRequest req; + + public Builder(String model, String input, String voice) { + this.req = new SpeechRequest(model, input, voice); + } + + public Builder responseFormat(String fmt) { + req.responseFormat = fmt; + return this; + } + + public Builder speed(double s) { + req.speed = s; + return this; + } + + public SpeechRequest build() { + return req; + } + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/SpeechResult.java b/bindings/jvm/src/main/java/io/aibridge/SpeechResult.java new file mode 100644 index 0000000..ad93032 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/SpeechResult.java @@ -0,0 +1,34 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonIgnoreProperties; +import com.fasterxml.jackson.annotation.JsonProperty; + +/** + * 文字转语音结果(对应 Rust {@code SpeechResult})。 + * + *

注意:Rust 侧 {@code audio_data} 被 {@code #[serde(skip)]},不参与 JSON 序列化, + * 二进制音频通过 FFI 的 {@code aibridge_bytes_t} 单独传递。本 POJO 仅承载 meta JSON + * 字段,二进制由 {@link Client#speech} 单独返回并封装到 {@link SpeechResultFull}。 + */ +@JsonIgnoreProperties(ignoreUnknown = true) +public class SpeechResult { + + /** 音频 URL(部分 Provider 返回) */ + @JsonProperty("audio_url") + public String audioUrl; + /** 音频 Base64 编码 */ + @JsonProperty("audio_base64") + public String audioBase64; + /** 音频 MIME 类型 */ + @JsonProperty("content_type") + public String contentType; + /** 音频格式(mp3/wav/opus 等) */ + public String format; + /** 估计音频时长(秒) */ + public Double duration; + /** 使用的模型 ID */ + public String model; + + public SpeechResult() { + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/SpeechResultFull.java b/bindings/jvm/src/main/java/io/aibridge/SpeechResultFull.java new file mode 100644 index 0000000..8c27724 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/SpeechResultFull.java @@ -0,0 +1,36 @@ +package io.aibridge; + +/** + * 文字转语音完整结果(meta + 二进制音频)。 + * + *

Rust 侧 {@code SpeechResult.audio_data} 不参与 serde,二进制通过 FFI 的 + * {@code aibridge_bytes_t} 单独传递。本类把 meta JSON 与二进制音频合并为单一返回, + * 便于调用方使用。 + */ +public class SpeechResultFull { + + /** meta 信息(content_type/format/duration/model 等) */ + private final SpeechResult meta; + /** 二进制音频数据(若 Provider 返回二进制;否则为空数组) */ + private final byte[] audioData; + + public SpeechResultFull(SpeechResult meta, byte[] audioData) { + this.meta = meta; + this.audioData = audioData == null ? new byte[0] : audioData; + } + + /** meta 信息 */ + public SpeechResult getMeta() { + return meta; + } + + /** 二进制音频数据 */ + public byte[] getAudioData() { + return audioData; + } + + /** 音频长度(便捷方法) */ + public int audioLength() { + return audioData.length; + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/TimeoutException.java b/bindings/jvm/src/main/java/io/aibridge/TimeoutException.java new file mode 100644 index 0000000..d61181a --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/TimeoutException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 超时错误(对应 code = "timeout_error") */ +public class TimeoutException extends AibridgeException { + public TimeoutException(String message, String details, boolean retryable) { + super(CODE_TIMEOUT, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/UnsupportedCapabilityException.java b/bindings/jvm/src/main/java/io/aibridge/UnsupportedCapabilityException.java new file mode 100644 index 0000000..971c749 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/UnsupportedCapabilityException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 不支持的能力(对应 code = "unsupported_capability") */ +public class UnsupportedCapabilityException extends AibridgeException { + public UnsupportedCapabilityException(String message, String details, boolean retryable) { + super(CODE_UNSUPPORTED_CAPABILITY, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/ValidationException.java b/bindings/jvm/src/main/java/io/aibridge/ValidationException.java new file mode 100644 index 0000000..151ee03 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/ValidationException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 参数校验错误(对应 code = "validation_error") */ +public class ValidationException extends AibridgeException { + public ValidationException(String message, String details, boolean retryable) { + super(CODE_VALIDATION, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/VoiceNotAvailableException.java b/bindings/jvm/src/main/java/io/aibridge/VoiceNotAvailableException.java new file mode 100644 index 0000000..350e16f --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/VoiceNotAvailableException.java @@ -0,0 +1,8 @@ +package io.aibridge; + +/** 音色不可用(对应 code = "voice_not_available") */ +public class VoiceNotAvailableException extends AibridgeException { + public VoiceNotAvailableException(String message, String details, boolean retryable) { + super(CODE_VOICE_NOT_AVAILABLE, message, details, retryable); + } +} diff --git a/bindings/jvm/src/main/java/io/aibridge/VoiceSpec.java b/bindings/jvm/src/main/java/io/aibridge/VoiceSpec.java new file mode 100644 index 0000000..7ba0329 --- /dev/null +++ b/bindings/jvm/src/main/java/io/aibridge/VoiceSpec.java @@ -0,0 +1,42 @@ +package io.aibridge; + +import com.fasterxml.jackson.annotation.JsonInclude; +import com.fasterxml.jackson.annotation.JsonProperty; + +import java.util.Arrays; +import java.util.List; + +/** + * 音色规格(对应 Rust {@code VoiceSpec},候选列表用于自动降级)。 + * + *

序列化形如 {@code {"voices":["alloy"]}}。 + */ +@JsonInclude(JsonInclude.Include.NON_NULL) +public class VoiceSpec { + + /** 音色列表(至少 1 个;多个时启用 fallback 降级) */ + public List voices; + + public VoiceSpec() { + } + + public VoiceSpec(List voices) { + this.voices = voices; + } + + /** 单个音色 */ + public static VoiceSpec single(String voice) { + return new VoiceSpec(Arrays.asList(voice)); + } + + /** 候选音色列表 */ + public static VoiceSpec multiple(String... voices) { + return new VoiceSpec(Arrays.asList(voices)); + } + + /** 主音色(列表第一个) */ + @JsonProperty("primary") + public String primary() { + return voices == null || voices.isEmpty() ? null : voices.get(0); + } +} From 2ecfb58c79f8dc10999fad167b1cf3c654008687 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 12:03:18 +0800 Subject: [PATCH 09/55] =?UTF-8?q?feat(aibridge-node):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?0.6=20napi-rs=20=E7=BB=91=E5=AE=9A=20+=20hello=20world?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-node/Cargo.toml | 6 +- crates/aibridge-node/index.d.ts | 160 ++++++++++ crates/aibridge-node/index.js | 316 +++++++++++++++++++ crates/aibridge-node/lib.js | 124 ++++++++ crates/aibridge-node/package-lock.json | 36 +++ crates/aibridge-node/package.json | 38 +++ crates/aibridge-node/src/lib.rs | 420 ++++++++++++++++++++++++- examples/hello_node.js | 88 ++++++ 8 files changed, 1184 insertions(+), 4 deletions(-) create mode 100644 crates/aibridge-node/index.d.ts create mode 100644 crates/aibridge-node/index.js create mode 100644 crates/aibridge-node/lib.js create mode 100644 crates/aibridge-node/package-lock.json create mode 100644 crates/aibridge-node/package.json create mode 100644 examples/hello_node.js diff --git a/crates/aibridge-node/Cargo.toml b/crates/aibridge-node/Cargo.toml index 5a73f2d..8ffefe1 100644 --- a/crates/aibridge-node/Cargo.toml +++ b/crates/aibridge-node/Cargo.toml @@ -2,7 +2,7 @@ name = "aibridge-node" version.workspace = true edition.workspace = true -rust-version.workspace = true +rust-version = "1.77" license.workspace = true authors.workspace = true repository.workspace = true @@ -13,9 +13,11 @@ crate-type = ["cdylib"] [dependencies] aibridge-core.workspace = true -napi = { version = "2", features = ["napi8", "tokio_rt"] } +# napi: 启用 tokio_rt(async fn → Promise)+ napi8(ThreadsafeFunction 等)+ serde-json(serde_json::Value 互转) +napi = { version = "2", features = ["napi8", "tokio_rt", "serde-json"] } napi-derive = "2" tokio.workspace = true +futures.workspace = true serde_json.workspace = true [build-dependencies] diff --git a/crates/aibridge-node/index.d.ts b/crates/aibridge-node/index.d.ts new file mode 100644 index 0000000..78e6f7b --- /dev/null +++ b/crates/aibridge-node/index.d.ts @@ -0,0 +1,160 @@ +/* tslint:disable */ +/* eslint-disable */ + +/* auto-generated by NAPI-RS */ + +/** 对话选项中的消息(JS 返回值形态) */ +export interface ChatChoiceJs { + /** 选项索引 */ + index: number + /** 生成的回复消息 */ + message: ChoiceMessageJs + /** 结束原因(stop / length / content_filter / tool_calls) */ + finishReason?: string +} +/** 完成结果中的消息 */ +export interface ChoiceMessageJs { + /** 角色(通常为 "assistant") */ + role: string + /** 消息内容 */ + content?: string +} +/** Token 使用统计 */ +export interface ChatUsageJs { + /** 提示词 token 数 */ + promptTokens: number + /** 完成回复 token 数 */ + completionTokens: number + /** 总 token 数 */ + totalTokens: number +} +/** 对话完成结果 */ +export interface ChatCompletionJs { + /** 响应 ID */ + id: string + /** 对象类型 */ + object: string + /** 创建时间戳(秒) */ + created: number + /** 使用的模型 */ + model: string + /** 回复选项列表 */ + choices: Array + /** Token 使用统计 */ + usage?: ChatUsageJs + /** 服务层级 */ + serviceTier?: string + /** 系统指纹 */ + systemFingerprint?: string +} +/** 流式增量消息 */ +export interface DeltaMessageJs { + /** 角色(首个块通常为 "assistant") */ + role?: string + /** 增量内容 */ + content?: string +} +/** 流式增量 */ +export interface ChatCompletionDeltaJs { + /** 增量索引 */ + index: number + /** 增量消息内容 */ + delta: DeltaMessageJs + /** 结束原因 */ + finishReason?: string +} +/** 流式对话块 */ +export interface ChatCompletionChunkJs { + /** 响应 ID */ + id: string + /** 对象类型 */ + object: string + /** 创建时间戳(秒) */ + created: number + /** 使用的模型 */ + model: string + /** 增量选项列表 */ + choices: Array + /** Token 使用统计(仅末尾块可能出现) */ + usage?: ChatUsageJs +} +/** 文字转语音结果 */ +export interface SpeechResultJs { + /** 音频二进制数据(JS Buffer) */ + audioData: Buffer + /** 音频 MIME 类型 */ + contentType: string + /** 音频格式(mp3/wav/opus 等) */ + format: string + /** 估计音频时长(秒) */ + duration?: number + /** 使用的模型 ID */ + model?: string + /** 音频 URL(部分 Provider 返回) */ + audioUrl?: string +} +/** + * AIBridge 统一客户端 + * + * 直连 aibridge-core 的 `Client`。所有异步方法返回 JS `Promise`。 + * + * 示例(JS): + * ```js + * const { Client } = require('aibridge'); + * const client = new Client('echo', {}); + * await client.start(); + * const resp = await client.chat({ model: 'echo-chat', messages: [{ role: 'user', content: 'hello' }] }); + * await client.close(); + * ``` + */ +export declare class Client { + /** + * 创建客户端 + * + * `provider` 为 Provider 类型(如 "echo"、"openai"、"agnes")。 + * `options` 为连接选项(api_key / base_url / timeout 等,可选字段)。 + */ + constructor(provider: string, options?: any | undefined | null) + /** 启动客户端(初始化适配器) */ + start(): Promise + /** 关闭客户端(释放资源) */ + close(): Promise + /** + * 文本对话 + * + * `request` 形如 `{ model, messages: [{ role, content }], temperature?, max_tokens? }`。 + * 透传不识别的字段到 core 的 `extra`(厂商特有参数)。 + */ + chat(request: any): Promise + /** + * 流式文本对话 + * + * 返回 `ChatStreamIterator`,支持 JS `for await...of` 迭代 `ChatCompletionChunk`。 + */ + chatStream(request: any): Promise + /** + * 文字转语音 + * + * `request` 形如 `{ model, input, voice, response_format?, speed? }`。 + * `voice` 可为字符串(单个音色)或字符串数组(候选列表,用于自动降级)。 + * 返回 `SpeechResultJs`,`audio_data` 为 JS `Buffer`。 + */ + speech(request: any): Promise +} +/** + * 流式对话迭代器 + * + * 由 `Client.chatStream()` 返回。通过 `[Symbol.asyncIterator]`(index.js 中安装) + * 支持 JS `for await...of` 语法迭代 `ChatCompletionChunk`。 + * + * 也可直接调用 `.next()` 方法手动迭代(返回 `null` 表示结束)。 + */ +export declare class ChatStreamIterator { + /** + * 拉取下一个 chunk + * + * 返回 `ChatCompletionChunk`,流结束时返回 `null`。 + * 流内部错误以 `napi::Error`(reject)形式抛出。 + */ + next(): Promise +} diff --git a/crates/aibridge-node/index.js b/crates/aibridge-node/index.js new file mode 100644 index 0000000..9f54547 --- /dev/null +++ b/crates/aibridge-node/index.js @@ -0,0 +1,316 @@ +/* tslint:disable */ +/* eslint-disable */ +/* prettier-ignore */ + +/* auto-generated by NAPI-RS */ + +const { existsSync, readFileSync } = require('fs') +const { join } = require('path') + +const { platform, arch } = process + +let nativeBinding = null +let localFileExisted = false +let loadError = null + +function isMusl() { + // For Node 10 + if (!process.report || typeof process.report.getReport !== 'function') { + try { + const lddPath = require('child_process').execSync('which ldd').toString().trim() + return readFileSync(lddPath, 'utf8').includes('musl') + } catch (e) { + return true + } + } else { + const { glibcVersionRuntime } = process.report.getReport().header + return !glibcVersionRuntime + } +} + +switch (platform) { + case 'android': + switch (arch) { + case 'arm64': + localFileExisted = existsSync(join(__dirname, 'aibridge.android-arm64.node')) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.android-arm64.node') + } else { + nativeBinding = require('aibridge-android-arm64') + } + } catch (e) { + loadError = e + } + break + case 'arm': + localFileExisted = existsSync(join(__dirname, 'aibridge.android-arm-eabi.node')) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.android-arm-eabi.node') + } else { + nativeBinding = require('aibridge-android-arm-eabi') + } + } catch (e) { + loadError = e + } + break + default: + throw new Error(`Unsupported architecture on Android ${arch}`) + } + break + case 'win32': + switch (arch) { + case 'x64': + localFileExisted = existsSync( + join(__dirname, 'aibridge.win32-x64-msvc.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.win32-x64-msvc.node') + } else { + nativeBinding = require('aibridge-win32-x64-msvc') + } + } catch (e) { + loadError = e + } + break + case 'ia32': + localFileExisted = existsSync( + join(__dirname, 'aibridge.win32-ia32-msvc.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.win32-ia32-msvc.node') + } else { + nativeBinding = require('aibridge-win32-ia32-msvc') + } + } catch (e) { + loadError = e + } + break + case 'arm64': + localFileExisted = existsSync( + join(__dirname, 'aibridge.win32-arm64-msvc.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.win32-arm64-msvc.node') + } else { + nativeBinding = require('aibridge-win32-arm64-msvc') + } + } catch (e) { + loadError = e + } + break + default: + throw new Error(`Unsupported architecture on Windows: ${arch}`) + } + break + case 'darwin': + localFileExisted = existsSync(join(__dirname, 'aibridge.darwin-universal.node')) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.darwin-universal.node') + } else { + nativeBinding = require('aibridge-darwin-universal') + } + break + } catch {} + switch (arch) { + case 'x64': + localFileExisted = existsSync(join(__dirname, 'aibridge.darwin-x64.node')) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.darwin-x64.node') + } else { + nativeBinding = require('aibridge-darwin-x64') + } + } catch (e) { + loadError = e + } + break + case 'arm64': + localFileExisted = existsSync( + join(__dirname, 'aibridge.darwin-arm64.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.darwin-arm64.node') + } else { + nativeBinding = require('aibridge-darwin-arm64') + } + } catch (e) { + loadError = e + } + break + default: + throw new Error(`Unsupported architecture on macOS: ${arch}`) + } + break + case 'freebsd': + if (arch !== 'x64') { + throw new Error(`Unsupported architecture on FreeBSD: ${arch}`) + } + localFileExisted = existsSync(join(__dirname, 'aibridge.freebsd-x64.node')) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.freebsd-x64.node') + } else { + nativeBinding = require('aibridge-freebsd-x64') + } + } catch (e) { + loadError = e + } + break + case 'linux': + switch (arch) { + case 'x64': + if (isMusl()) { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-x64-musl.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-x64-musl.node') + } else { + nativeBinding = require('aibridge-linux-x64-musl') + } + } catch (e) { + loadError = e + } + } else { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-x64-gnu.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-x64-gnu.node') + } else { + nativeBinding = require('aibridge-linux-x64-gnu') + } + } catch (e) { + loadError = e + } + } + break + case 'arm64': + if (isMusl()) { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-arm64-musl.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-arm64-musl.node') + } else { + nativeBinding = require('aibridge-linux-arm64-musl') + } + } catch (e) { + loadError = e + } + } else { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-arm64-gnu.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-arm64-gnu.node') + } else { + nativeBinding = require('aibridge-linux-arm64-gnu') + } + } catch (e) { + loadError = e + } + } + break + case 'arm': + if (isMusl()) { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-arm-musleabihf.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-arm-musleabihf.node') + } else { + nativeBinding = require('aibridge-linux-arm-musleabihf') + } + } catch (e) { + loadError = e + } + } else { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-arm-gnueabihf.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-arm-gnueabihf.node') + } else { + nativeBinding = require('aibridge-linux-arm-gnueabihf') + } + } catch (e) { + loadError = e + } + } + break + case 'riscv64': + if (isMusl()) { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-riscv64-musl.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-riscv64-musl.node') + } else { + nativeBinding = require('aibridge-linux-riscv64-musl') + } + } catch (e) { + loadError = e + } + } else { + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-riscv64-gnu.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-riscv64-gnu.node') + } else { + nativeBinding = require('aibridge-linux-riscv64-gnu') + } + } catch (e) { + loadError = e + } + } + break + case 's390x': + localFileExisted = existsSync( + join(__dirname, 'aibridge.linux-s390x-gnu.node') + ) + try { + if (localFileExisted) { + nativeBinding = require('./aibridge.linux-s390x-gnu.node') + } else { + nativeBinding = require('aibridge-linux-s390x-gnu') + } + } catch (e) { + loadError = e + } + break + default: + throw new Error(`Unsupported architecture on Linux: ${arch}`) + } + break + default: + throw new Error(`Unsupported OS: ${platform}, architecture: ${arch}`) +} + +if (!nativeBinding) { + if (loadError) { + throw loadError + } + throw new Error(`Failed to load native binding`) +} + +const { Client, ChatStreamIterator } = nativeBinding + +module.exports.Client = Client +module.exports.ChatStreamIterator = ChatStreamIterator diff --git a/crates/aibridge-node/lib.js b/crates/aibridge-node/lib.js new file mode 100644 index 0000000..f843e87 --- /dev/null +++ b/crates/aibridge-node/lib.js @@ -0,0 +1,124 @@ +'use strict'; + +// AIBridge Node.js 绑定入口(含 JS 侧包装) +// +// napi build 会自动生成 index.js(纯 native re-export),因此本文件作为 +// package.json 的 main 入口,负责: +// 1. 加载 native 模块 +// 2. 包装 Client 的所有方法与构造函数,统一为抛出的 Error 解析出 `.code` 属性 +// (Rust 侧 map_error 将 AibridgeError 编码为 `[code] message`) +// 3. 为 ChatStreamIterator.prototype 安装 [Symbol.asyncIterator],支持 `for await...of` + +const native = require('./index.js'); +const NativeClient = native.Client; +const { ChatStreamIterator } = native; + +/** + * 错误码前缀正则:匹配 `[code] message` 格式 + */ +const ERROR_CODE_RE = /^\[([a-z_]+)\]\s*(.*)$/s; + +/** + * 包装 Error,解析出 `.code` 属性 + * + * Rust 侧 map_error 将 AibridgeError 编码为 `[code] message`, + * 此处解析出 code 并挂到 Error.code 属性上;无前缀的视为 unknown_error。 + * @param {Error} err - 原始错误 + * @returns {Error} 带 .code 属性的错误 + */ +function withCode(err) { + if (err && typeof err.message === 'string') { + const m = err.message.match(ERROR_CODE_RE); + if (m) { + err.code = m[1]; + err.message = m[2]; + } else if (!err.code) { + err.code = 'unknown_error'; + } + } + return err; +} + +/** + * 包装一个返回 Promise 的方法,reject 时解析 .code + */ +function wrapAsync(fn) { + return function (...args) { + const p = fn.apply(this, args); + if (p && typeof p.then === 'function') { + return p.catch((err) => Promise.reject(withCode(err))); + } + return p; + }; +} + +/** + * AIBridge 统一客户端(JS 包装层) + * + * 代理原生 Client,构造与方法调用的错误统一经 withCode 解析出 `.code`。 + * 其余行为与原生 Client 完全一致。 + */ +class Client { + constructor(provider, options) { + try { + this._native = new NativeClient(provider, options); + } catch (err) { + throw withCode(err); + } + } + + start() { + return wrapAsync(() => this._native.start()).call(this); + } + + close() { + return wrapAsync(() => this._native.close()).call(this); + } + + chat(request) { + return wrapAsync(() => this._native.chat(request)).call(this); + } + + speech(request) { + return wrapAsync(() => this._native.speech(request)).call(this); + } + + chatStream(request) { + // chatStream 返回 Promise,reject 时需解析 code + return this._native.chatStream(request).catch((err) => + Promise.reject(withCode(err)) + ); + } +} + +/** + * 为 ChatStreamIterator.prototype 安装 [Symbol.asyncIterator] + * + * napi 2 的 Generator trait 是同步的,无法桥接异步 stream,因此通过 next() 方法 + * (返回 Promise)手动实现 asyncIterator 协议。 + */ +if (ChatStreamIterator && !ChatStreamIterator.prototype[Symbol.asyncIterator]) { + ChatStreamIterator.prototype[Symbol.asyncIterator] = function asyncIterator() { + const self = this; + return { + async next() { + try { + const chunk = await self.next(); + // chunk === null 表示流结束 + if (chunk === null || chunk === undefined) { + return { value: undefined, done: true }; + } + return { value: chunk, done: false }; + } catch (err) { + // 流内部错误,解析 code 后抛出 + throw withCode(err); + } + }, + }; + }; +} + +module.exports = { Client, ChatStreamIterator }; +module.exports.default = module.exports; +module.exports.Client = Client; +module.exports.ChatStreamIterator = ChatStreamIterator; diff --git a/crates/aibridge-node/package-lock.json b/crates/aibridge-node/package-lock.json new file mode 100644 index 0000000..536dd34 --- /dev/null +++ b/crates/aibridge-node/package-lock.json @@ -0,0 +1,36 @@ +{ + "name": "aibridge", + "version": "2.0.0-alpha.1", + "lockfileVersion": 3, + "requires": true, + "packages": { + "": { + "name": "aibridge", + "version": "2.0.0-alpha.1", + "license": "MIT", + "devDependencies": { + "@napi-rs/cli": "^2.18.4" + }, + "engines": { + "node": ">= 16" + } + }, + "node_modules/@napi-rs/cli": { + "version": "2.18.4", + "resolved": "https://registry.npmjs.org/@napi-rs/cli/-/cli-2.18.4.tgz", + "integrity": "sha512-SgJeA4df9DE2iAEpr3M2H0OKl/yjtg1BnRI5/JyowS71tUWhrfSu2LT0V3vlHET+g1hBVlrO60PmEXwUEKp8Mg==", + "dev": true, + "license": "MIT", + "bin": { + "napi": "scripts/index.js" + }, + "engines": { + "node": ">= 10" + }, + "funding": { + "type": "github", + "url": "https://github.com/sponsors/Brooooooklyn" + } + } + } +} diff --git a/crates/aibridge-node/package.json b/crates/aibridge-node/package.json new file mode 100644 index 0000000..191863c --- /dev/null +++ b/crates/aibridge-node/package.json @@ -0,0 +1,38 @@ +{ + "name": "aibridge", + "version": "2.0.0-alpha.1", + "description": "AIBridge Node.js 绑定(napi-rs,直连 aibridge-core,原生 Promise)", + "main": "lib.js", + "types": "index.d.ts", + "license": "MIT", + "author": "WingkySky ", + "repository": { + "type": "git", + "url": "https://github.com/WingkySky/aibridge", + "directory": "crates/aibridge-node" + }, + "keywords": [ + "aibridge", + "ai", + "napi-rs", + "llm", + "chat", + "tts" + ], + "engines": { + "node": ">= 16" + }, + "napi": { + "name": "aibridge", + "triples": {} + }, + "scripts": { + "build": "napi build --platform", + "build:debug": "napi build --platform -d", + "build:release": "napi build --platform --release", + "prepublishOnly": "napi prepublish -t npm" + }, + "devDependencies": { + "@napi-rs/cli": "^2.18.4" + } +} diff --git a/crates/aibridge-node/src/lib.rs b/crates/aibridge-node/src/lib.rs index ddf4daa..37f924d 100644 --- a/crates/aibridge-node/src/lib.rs +++ b/crates/aibridge-node/src/lib.rs @@ -2,6 +2,422 @@ //! //! 直连 aibridge-core,原生 Promise / AsyncIterable 流式。 //! 由 napi-rs 构建为 npm 包 `aibridge`。 -//! 阶段 0.6 填充 Client/chat/流式/错误映射。 +//! +//! 设计要点(对应设计文档第 8 节 JS 桥接): +//! - `#[napi] async fn` 自动桥接为 JS `Promise`(基于 napi 全局 tokio runtime) +//! - 复杂请求参数通过 `serde_json::Value` 中转(JS Object → Value → core serde struct) +//! - 返回值用 `#[napi(object)]` struct,JS 拿到原生对象,类型清晰 +//! - 流式:`chatStream` 返回 `ChatStreamIterator`,内部 spawn tokio task 消费 core 的 +//! `BoxStream`,通过 `tokio::sync::mpsc` channel 推送 chunk;JS 侧通过 `[Symbol.asyncIterator]` +//! 支持 `for await...of`(见 index.js 的包装) +//! - 错误映射:`AibridgeError` → `napi::Error`,reason 编码为 `[code] message`, +//! JS 侧 index.js 解析出 `.code` 属性(napi 的 Error code 字段无法承载自定义业务 code) + +use std::sync::Arc; + +use futures::stream::StreamExt; +use napi::bindgen_prelude::*; +use napi_derive::napi; +use serde_json::Value; +use tokio::sync::{mpsc, Mutex}; + +use aibridge_core::client::Client as CoreClient; +use aibridge_core::config::ClientOptions; +use aibridge_core::error::AibridgeError; +use aibridge_core::model::audio::SpeechRequest; +use aibridge_core::model::chat::{ChatCompletion, ChatCompletionChunk, ChatRequest}; + +// ────────────────────────────────────────────────────────────────────────── +// 错误映射 +// ────────────────────────────────────────────────────────────────────────── + +/// 将 `AibridgeError` 映射为 `napi::Error` +/// +/// reason 编码为 `[code] message` 形式,便于 JS 侧解析 `.code` 属性。 +/// status 统一用 `GenericFailure`(napi 的 Status 枚举无法承载业务 code)。 +fn map_error(err: AibridgeError) -> Error { + let code = err.code(); + let message = err.to_string(); + Error::new( + Status::GenericFailure, + format!("[{code}] {message}"), + ) +} + +// ────────────────────────────────────────────────────────────────────────── +// JS 友好的返回数据模型(#[napi(object)]) +// ────────────────────────────────────────────────────────────────────────── + +/// 对话选项中的消息(JS 返回值形态) +#[napi(object)] +pub struct ChatChoiceJs { + /// 选项索引 + pub index: u32, + /// 生成的回复消息 + pub message: ChoiceMessageJs, + /// 结束原因(stop / length / content_filter / tool_calls) + pub finish_reason: Option, +} + +/// 完成结果中的消息 +#[napi(object)] +pub struct ChoiceMessageJs { + /// 角色(通常为 "assistant") + pub role: String, + /// 消息内容 + pub content: Option, +} + +/// Token 使用统计 +#[napi(object)] +pub struct ChatUsageJs { + /// 提示词 token 数 + pub prompt_tokens: i64, + /// 完成回复 token 数 + pub completion_tokens: i64, + /// 总 token 数 + pub total_tokens: i64, +} + +/// 对话完成结果 +#[napi(object)] +pub struct ChatCompletionJs { + /// 响应 ID + pub id: String, + /// 对象类型 + pub object: String, + /// 创建时间戳(秒) + pub created: i64, + /// 使用的模型 + pub model: String, + /// 回复选项列表 + pub choices: Vec, + /// Token 使用统计 + pub usage: Option, + /// 服务层级 + pub service_tier: Option, + /// 系统指纹 + pub system_fingerprint: Option, +} + +/// 流式增量消息 +#[napi(object)] +pub struct DeltaMessageJs { + /// 角色(首个块通常为 "assistant") + pub role: Option, + /// 增量内容 + pub content: Option, +} + +/// 流式增量 +#[napi(object)] +pub struct ChatCompletionDeltaJs { + /// 增量索引 + pub index: u32, + /// 增量消息内容 + pub delta: DeltaMessageJs, + /// 结束原因 + pub finish_reason: Option, +} + +/// 流式对话块 +#[napi(object)] +pub struct ChatCompletionChunkJs { + /// 响应 ID + pub id: String, + /// 对象类型 + pub object: String, + /// 创建时间戳(秒) + pub created: i64, + /// 使用的模型 + pub model: String, + /// 增量选项列表 + pub choices: Vec, + /// Token 使用统计(仅末尾块可能出现) + pub usage: Option, +} + +/// 文字转语音结果 +#[napi(object)] +pub struct SpeechResultJs { + /// 音频二进制数据(JS Buffer) + pub audio_data: Buffer, + /// 音频 MIME 类型 + pub content_type: String, + /// 音频格式(mp3/wav/opus 等) + pub format: String, + /// 估计音频时长(秒) + pub duration: Option, + /// 使用的模型 ID + pub model: Option, + /// 音频 URL(部分 Provider 返回) + pub audio_url: Option, +} + +// ────────────────────────────────────────────────────────────────────────── +// core → JS 数据模型转换 +// ────────────────────────────────────────────────────────────────────────── + +/// 将 core 的 `ChatCompletion` 转为 JS 友好结构 +fn to_chat_completion_js(c: ChatCompletion) -> ChatCompletionJs { + ChatCompletionJs { + id: c.id, + object: c.object, + created: c.created as i64, + model: c.model, + choices: c + .choices + .into_iter() + .map(|ch| ChatChoiceJs { + index: ch.index, + message: ChoiceMessageJs { + role: ch.message.role, + content: ch.message.content, + }, + finish_reason: ch.finish_reason, + }) + .collect(), + usage: c.usage.map(|u| ChatUsageJs { + prompt_tokens: u.prompt_tokens as i64, + completion_tokens: u.completion_tokens as i64, + total_tokens: u.total_tokens as i64, + }), + service_tier: c.service_tier, + system_fingerprint: c.system_fingerprint, + } +} + +/// 将 core 的 `ChatCompletionChunk` 转为 JS 友好结构 +fn to_chat_chunk_js(c: ChatCompletionChunk) -> ChatCompletionChunkJs { + ChatCompletionChunkJs { + id: c.id, + object: c.object, + created: c.created as i64, + model: c.model, + choices: c + .choices + .into_iter() + .map(|d| ChatCompletionDeltaJs { + index: d.index, + delta: DeltaMessageJs { + role: d.delta.role, + content: d.delta.content, + }, + finish_reason: d.finish_reason, + }) + .collect(), + usage: c.usage.map(|u| ChatUsageJs { + prompt_tokens: u.prompt_tokens as i64, + completion_tokens: u.completion_tokens as i64, + total_tokens: u.total_tokens as i64, + }), + } +} + +// ────────────────────────────────────────────────────────────────────────── +// 统一客户端(napi 类) +// ────────────────────────────────────────────────────────────────────────── + +/// AIBridge 统一客户端 +/// +/// 直连 aibridge-core 的 `Client`。所有异步方法返回 JS `Promise`。 +/// +/// 示例(JS): +/// ```js +/// const { Client } = require('aibridge'); +/// const client = new Client('echo', {}); +/// await client.start(); +/// const resp = await client.chat({ model: 'echo-chat', messages: [{ role: 'user', content: 'hello' }] }); +/// await client.close(); +/// ``` +#[napi] +pub struct Client { + /// core 客户端,用 `Arc` 包裹以支持 `&self` async 方法 + /// (napi async fn 不允许 `&mut self`,需内部可变性) + inner: Arc>, +} + +#[napi] +impl Client { + /// 创建客户端 + /// + /// `provider` 为 Provider 类型(如 "echo"、"openai"、"agnes")。 + /// `options` 为连接选项(api_key / base_url / timeout 等,可选字段)。 + #[napi(constructor)] + pub fn new(provider: String, options: Option) -> Result { + // 将 JS options 对象转为 ClientOptions(经 serde_json 中转) + let opts: ClientOptions = match options { + Some(v) => serde_json::from_value(v) + .map_err(|e| Error::new(Status::InvalidArg, format!("options 解析失败: {e}")))?, + None => ClientOptions::default(), + }; + + let core = CoreClient::new(&provider, opts).map_err(map_error)?; + Ok(Self { + inner: Arc::new(Mutex::new(core)), + }) + } + + /// 启动客户端(初始化适配器) + #[napi] + pub async fn start(&self) -> Result<()> { + let mut guard = self.inner.lock().await; + guard.start().await.map_err(map_error) + } + + /// 关闭客户端(释放资源) + #[napi] + pub async fn close(&self) -> Result<()> { + let mut guard = self.inner.lock().await; + guard.close().await.map_err(map_error) + } + + /// 文本对话 + /// + /// `request` 形如 `{ model, messages: [{ role, content }], temperature?, max_tokens? }`。 + /// 透传不识别的字段到 core 的 `extra`(厂商特有参数)。 + #[napi] + pub async fn chat(&self, request: Value) -> Result { + let req: ChatRequest = serde_json::from_value(request) + .map_err(|e| Error::new(Status::InvalidArg, format!("request 解析失败: {e}")))?; + let guard = self.inner.lock().await; + let resp = guard.chat(req).await.map_err(map_error)?; + Ok(to_chat_completion_js(resp)) + } + + /// 流式文本对话 + /// + /// 返回 `ChatStreamIterator`,支持 JS `for await...of` 迭代 `ChatCompletionChunk`。 + #[napi] + pub async fn chat_stream(&self, request: Value) -> Result { + let req: ChatRequest = serde_json::from_value(request) + .map_err(|e| Error::new(Status::InvalidArg, format!("request 解析失败: {e}")))?; + + // 在 napi 全局 tokio runtime 上获取流(需要在持锁期间拿到 stream) + let stream = { + let guard = self.inner.lock().await; + guard.chat_stream(req).await.map_err(map_error)? + }; + + // 用 mpsc channel 桥接 BoxStream:spawn 一个 task 消费 stream 并推送 chunk + let (tx, rx) = mpsc::channel::>(16); + // spawn 到 napi 全局 tokio runtime(fire-and-forget,task 自行消费完毕后 drop tx) + spawn(async move { + let mut stream = stream; + while let Some(item) = stream.next().await { + let chunk_result = item.map(to_chat_chunk_js).map_err(map_error); + // 接收方关闭则结束 + if tx.send(chunk_result).await.is_err() { + break; + } + } + // stream 结束(正常或错误已推送),drop tx 让 rx 收到 None 表示完成 + drop(tx); + }); + + Ok(ChatStreamIterator { + rx: Arc::new(Mutex::new(rx)), + }) + } + + /// 文字转语音 + /// + /// `request` 形如 `{ model, input, voice, response_format?, speed? }`。 + /// `voice` 可为字符串(单个音色)或字符串数组(候选列表,用于自动降级)。 + /// 返回 `SpeechResultJs`,`audio_data` 为 JS `Buffer`。 + #[napi] + pub async fn speech(&self, request: Value) -> Result { + // 将 voice 字段归一化为 VoiceSpec 结构(core 期望 { voices: [...] }) + let mut req_value = request; + if let Some(obj) = req_value.as_object_mut() { + if let Some(voice) = obj.remove("voice") { + let voices = match voice { + // 字符串 → [voice] + Value::String(s) => vec![s], + // 字符串数组 → 原样 + Value::Array(arr) => arr + .into_iter() + .filter_map(|v| v.as_str().map(str::to_owned)) + .collect(), + // 已是 { voices: [...] } 对象 → 直接用 + Value::Object(map) if map.contains_key("voices") => { + match map.get("voices").cloned() { + Some(Value::Array(arr)) => arr + .into_iter() + .filter_map(|v| v.as_str().map(str::to_owned)) + .collect(), + _ => { + return Err(Error::new( + Status::InvalidArg, + "voice.voices 必须为字符串数组", + )) + } + } + } + _ => { + return Err(Error::new( + Status::InvalidArg, + "voice 必须为字符串、字符串数组或 { voices: [...] } 对象", + )) + } + }; + obj.insert( + "voice".into(), + serde_json::json!({ "voices": voices }), + ); + } + } + + let req: SpeechRequest = serde_json::from_value(req_value) + .map_err(|e| Error::new(Status::InvalidArg, format!("request 解析失败: {e}")))?; + let guard = self.inner.lock().await; + let resp = guard.speech(req).await.map_err(map_error)?; + + // 优先 audio_data,其次解码 audio_base64,最后空 Buffer + let bytes = resp.get_audio_bytes().unwrap_or_default(); + Ok(SpeechResultJs { + audio_data: Buffer::from(bytes), + content_type: resp.content_type, + format: resp.format, + duration: resp.duration, + model: resp.model, + audio_url: resp.audio_url, + }) + } +} + +// ────────────────────────────────────────────────────────────────────────── +// 流式迭代器(napi 类) +// ────────────────────────────────────────────────────────────────────────── + +/// 流式对话迭代器 +/// +/// 由 `Client.chatStream()` 返回。通过 `[Symbol.asyncIterator]`(index.js 中安装) +/// 支持 JS `for await...of` 语法迭代 `ChatCompletionChunk`。 +/// +/// 也可直接调用 `.next()` 方法手动迭代(返回 `null` 表示结束)。 +#[napi] +pub struct ChatStreamIterator { + /// chunk 接收端,用 `Arc` 包裹以支持 `&self` async 方法 + rx: Arc>>>, +} -// 阶段 0.6 填充:#[napi] 标注的 Client 及方法 +#[napi] +impl ChatStreamIterator { + /// 拉取下一个 chunk + /// + /// 返回 `ChatCompletionChunk`,流结束时返回 `null`。 + /// 流内部错误以 `napi::Error`(reject)形式抛出。 + #[napi] + pub async fn next(&self) -> Result> { + let mut guard = self.rx.lock().await; + match guard.recv().await { + // 收到正常 chunk + Some(Ok(chunk)) => Ok(Some(chunk)), + // 接收端收到错误 chunk + Some(Err(e)) => Err(e), + // channel 关闭,流结束 + None => Ok(None), + } + } +} diff --git a/examples/hello_node.js b/examples/hello_node.js new file mode 100644 index 0000000..a175da0 --- /dev/null +++ b/examples/hello_node.js @@ -0,0 +1,88 @@ +'use strict'; + +// AIBridge Node.js 绑定 hello world +// +// 使用 echo mock 适配器(免认证)端到端验证: +// 1. chat:回显用户消息,期望 choices[0].message.content === "hello [echo]" +// 2. chatStream:流式回显,期望 3 个 chunk +// 3. speech:返回固定音频,期望 audioData.length === 15 +// +// 运行:node examples/hello_node.js +// 前置:在 crates/aibridge-node 下执行 npm install && npm run build(或 cargo build -p aibridge-node) + +const { Client } = require('../crates/aibridge-node'); + +async function main() { + // 1. 创建客户端(echo 免认证,无需 api_key) + const client = new Client('echo', {}); + await client.start(); + console.log('[hello_node] 客户端已启动'); + + try { + // 2. chat:非流式对话 + const resp = await client.chat({ + model: 'echo-chat', + messages: [{ role: 'user', content: 'hello' }], + }); + const content = resp.choices[0].message.content; + console.log('[hello_node] chat 结果: %j', content); + if (content !== 'hello [echo]') { + throw new Error(`chat 断言失败:期望 "hello [echo]",实际 "${content}"`); + } + + // 3. chatStream:流式对话,for await 迭代 chunk + const stream = await client.chatStream({ + model: 'echo-chat', + messages: [{ role: 'user', content: 'hello' }], + }); + let chunkCount = 0; + let assembled = ''; + for await (const chunk of stream) { + chunkCount += 1; + const delta = chunk.choices[0]?.delta; + if (delta?.role) { + console.log('[hello_node] chunk %d role=%j', chunkCount, delta.role); + } + if (delta?.content) { + assembled += delta.content; + console.log('[hello_node] chunk %d content=%j', chunkCount, delta.content); + } + } + console.log('[hello_node] chatStream 共 %d 个 chunk,拼接内容=%j', chunkCount, assembled); + if (chunkCount !== 3) { + throw new Error(`chatStream 断言失败:期望 3 个 chunk,实际 ${chunkCount}`); + } + if (assembled !== 'hello [echo]') { + throw new Error(`chatStream 拼接断言失败:期望 "hello [echo]",实际 "${assembled}"`); + } + + // 4. speech:文字转语音,返回 Buffer + const speech = await client.speech({ + model: 'echo-tts', + input: 'hello', + voice: 'alloy', + }); + console.log( + '[hello_node] speech 结果: audioData.length=%d, format=%j, contentType=%j', + speech.audioData.length, + speech.format, + speech.contentType + ); + if (speech.audioData.length !== 15) { + throw new Error( + `speech 断言失败:期望 audioData.length === 15,实际 ${speech.audioData.length}` + ); + } + } finally { + // 5. 关闭客户端 + await client.close(); + console.log('[hello_node] 客户端已关闭'); + } + + console.log('[hello_node] 全部断言通过 ✓'); +} + +main().catch((err) => { + console.error('[hello_node] 失败:', err); + process.exit(1); +}); From 268cd49896d0ab64966176cc0669401480e7a086 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 12:07:32 +0800 Subject: [PATCH 10/55] =?UTF-8?q?feat(aibridge-python):=20=E9=98=B6?= =?UTF-8?q?=E6=AE=B50.6=20PyO3=20=E7=BB=91=E5=AE=9A=20+=20hello=20world=20?= =?UTF-8?q?=E7=AE=A1=E7=BA=BF=E9=AA=8C=E8=AF=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-python/Cargo.toml | 16 +- crates/aibridge-python/pyproject.toml | 18 + crates/aibridge-python/src/lib.rs | 890 +++++++++++++++++++++++++- examples/hello_python.py | 85 +++ 4 files changed, 1002 insertions(+), 7 deletions(-) create mode 100644 crates/aibridge-python/pyproject.toml create mode 100644 examples/hello_python.py diff --git a/crates/aibridge-python/Cargo.toml b/crates/aibridge-python/Cargo.toml index d23cdc2..dff497f 100644 --- a/crates/aibridge-python/Cargo.toml +++ b/crates/aibridge-python/Cargo.toml @@ -14,7 +14,17 @@ name = "aibridge" [dependencies] aibridge-core.workspace = true -# PyO3 0.22+ 原生支持 async fn,不再需要 pyo3-asyncio -pyo3 = { version = "0.28", features = ["extension-module", "abi3-py310"] } -tokio.workspace = true +# PyO3 0.28 原生支持 async fn(需 experimental-async feature),不再需要 pyo3-asyncio。 +# +# `extension-module` feature 不在此处启用:它会让 `cargo build` 因不链接 Python +# 动态库而失败(扩展模块符号由 Python 进程运行时提供)。改由 pyproject.toml 的 +# `[tool.maturin] features = ["pyo3/extension-module"]` 在 maturin 构建时启用, +# 这样 `cargo build -p aibridge-python` 能正常链接 Python 框架,maturin 产物正确。 +pyo3 = { version = "0.28", features = [ + "abi3-py310", + "experimental-async", +] } +tokio = { workspace = true, features = ["rt-multi-thread", "rt", "sync"] } +futures.workspace = true serde_json.workspace = true +once_cell.workspace = true diff --git a/crates/aibridge-python/pyproject.toml b/crates/aibridge-python/pyproject.toml new file mode 100644 index 0000000..fe6d650 --- /dev/null +++ b/crates/aibridge-python/pyproject.toml @@ -0,0 +1,18 @@ +[build-system] +requires = ["maturin>=1.0,<2.0"] +build-backend = "maturin" + +[project] +name = "aibridge" +description = "AIBridge Python 绑定(PyO3,直连 aibridge-core,原生 asyncio)" +requires-python = ">=3.10" +classifiers = [ + "Programming Language :: Python :: 3", + "Programming Language :: Rust", + "License :: OSI Approved :: MIT License", + "Operating System :: OS Independent", +] +dynamic = ["version"] + +[tool.maturin] +features = ["pyo3/extension-module"] diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs index c1f498c..e9df520 100644 --- a/crates/aibridge-python/src/lib.rs +++ b/crates/aibridge-python/src/lib.rs @@ -2,13 +2,895 @@ //! //! 直连 aibridge-core,原生 asyncio 协程与 AsyncIterator 流式。 //! 由 maturin 构建为 PyPI 包 `aibridge`。 -//! 阶段 0.6 填充 Client/chat/流式/错误映射。 +//! +//! 阶段 0.6 实现:Client / chat / speech / chat_stream / 错误映射 / 数据模型。 +//! +//! 架构要点: +//! - 全局多线程 tokio runtime:core 的 async future(含 reqwest 等真实 IO)spawn 到 +//! tokio 上执行,PyO3 协程通过 await JoinHandle 拿结果。echo adapter 无网络也走 +//! 同一路径,保持一致。 +//! - `PyClient` 持有 `Arc>`,支持 `start`/`close` 可变操作。 +//! - 流式:`chat_stream` 把 core `ChatStream` 通过 `tokio::sync::Mutex` 封装进 +//! `ChatStreamIterator`,`__anext__` 同步返回 awaitable(PyO3 0.28 slot 不支持 +//! async fn),awaitable 内抛 `StopIteration(chunk)`/`StopAsyncIteration`。 + +// PyO3 0.28 的 `#[pymethods]` 宏展开会用到 Rust 1.77+ 稳定的语法(如 let 链), +// 与 workspace MSRV 1.75 冲突。该 lint 针对宏生成代码,非手写代码,故整体允许。 +#![allow(clippy::incompatible_msrv)] + +use std::sync::Arc; +use futures::StreamExt; use pyo3::prelude::*; +use pyo3::types::PyBytes; +use tokio::sync::Mutex; + +use aibridge_core::adapter::ChatStream as CoreChatStream; +use aibridge_core::client::Client as CoreClient; +use aibridge_core::config::ClientOptions as CoreClientOptions; +use aibridge_core::error::AibridgeError as CoreAibridgeError; +use aibridge_core::model::audio::{ + SpeechRequest as CoreSpeechRequest, SpeechResult as CoreSpeechResult, +}; +use aibridge_core::model::chat::{ + ChatCompletion as CoreChatCompletion, ChatCompletionChunk as CoreChatCompletionChunk, + ChatMessage as CoreChatMessage, ChatRequest as CoreChatRequest, +}; + +// =========================================================================== +// 全局 tokio runtime +// =========================================================================== + +/// 全局多线程 tokio runtime +/// +/// 用 `once_cell::sync::Lazy` 在首次访问时初始化。core 的 async future(含 +/// reqwest 等真实 IO)spawn 到此 runtime 上执行,PyO3 协程通过 await JoinHandle +/// 取回结果。多线程 runtime 保证真实 adapter 的并发能力。 +static RUNTIME: once_cell::sync::Lazy = + once_cell::sync::Lazy::new(|| { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + .expect("初始化 tokio runtime 失败") + }); + +// =========================================================================== +// 错误映射 +// =========================================================================== + +// 异常类层级(设计文档 9.3 节): +// AibridgeError (基类,继承 PyException) +// ├── AuthenticationError +// ├── RateLimitError +// ├── ValidationError +// ├── ModelNotFoundError +// ├── APIError +// ├── NetworkError +// ├── TimeoutError +// ├── UnsupportedCapabilityError +// ├── ProviderNotFoundError +// ├── VoiceNotAvailableError +// └── ServiceUnavailableError +// +// 子类名与 Python v1 (agn-sdk) 保持一致,便于迁移(v1 `AGNError` → v2 `AibridgeError`)。 +// 用 `create_exception!` 宏生成,构造方式 `XxxError::new_err(message)`。 +use pyo3::create_exception; + +create_exception!( + aibridge, + AibridgeError, + pyo3::exceptions::PyException, + "AIBridge SDK 错误基类" +); +create_exception!( + aibridge, + AuthenticationError, + AibridgeError, + "认证失败(API Key 无效/过期/无权限)" +); +create_exception!( + aibridge, + RateLimitError, + AibridgeError, + "请求频率超过限制" +); +create_exception!( + aibridge, + ValidationError, + AibridgeError, + "请求参数校验错误" +); +create_exception!( + aibridge, + ModelNotFoundError, + AibridgeError, + "请求的模型不存在" +); +create_exception!( + aibridge, + APIError, + AibridgeError, + "Provider API 调用错误" +); +create_exception!( + aibridge, + NetworkError, + AibridgeError, + "网络错误" +); +create_exception!( + aibridge, + TimeoutError, + AibridgeError, + "请求超时" +); +create_exception!( + aibridge, + UnsupportedCapabilityError, + AibridgeError, + "Provider 不支持该能力" +); +create_exception!( + aibridge, + ProviderNotFoundError, + AibridgeError, + "Provider 不存在" +); +create_exception!( + aibridge, + VoiceNotAvailableError, + AibridgeError, + "音色不可用" +); +create_exception!( + aibridge, + ServiceUnavailableError, + AibridgeError, + "服务暂时不可用(可重试)" +); + +/// 将 core `AibridgeError` 映射为对应的 Python 异常 `PyErr` +/// +/// 对齐设计文档 9.3 节错误映射表。消息格式为 `[code] message`,便于调用方 +/// 获取稳定标识码(core `AibridgeError::code()`)。 +fn map_error(err: CoreAibridgeError) -> PyErr { + let code = err.code(); + let message = format!("[{code}] {err}"); + match err { + CoreAibridgeError::Authentication { .. } => AuthenticationError::new_err(message), + CoreAibridgeError::RateLimit { .. } => RateLimitError::new_err(message), + CoreAibridgeError::Validation { .. } => ValidationError::new_err(message), + CoreAibridgeError::ModelNotFound { .. } => ModelNotFoundError::new_err(message), + CoreAibridgeError::Api { .. } => APIError::new_err(message), + CoreAibridgeError::Network(_) => NetworkError::new_err(message), + CoreAibridgeError::Timeout => TimeoutError::new_err(message), + CoreAibridgeError::UnsupportedCapability { .. } => { + UnsupportedCapabilityError::new_err(message) + } + CoreAibridgeError::ProviderNotFound { .. } => ProviderNotFoundError::new_err(message), + CoreAibridgeError::VoiceNotAvailable { .. } => VoiceNotAvailableError::new_err(message), + CoreAibridgeError::ServiceUnavailable { .. } => ServiceUnavailableError::new_err(message), + } +} + +// =========================================================================== +// 数据模型 +// =========================================================================== + +/// 对话消息 +/// +/// 对应 core `ChatMessage`。Python 侧用 dict 构造(`{"role": "user", "content": "..."}`), +/// 仅支持 chat hello world 所需的 system/user/assistant 文本消息。 +#[pyclass(from_py_object)] +#[derive(Debug, Clone)] +struct ChatMessage { + /// 角色(system / user / assistant / tool) + role: String, + /// 文本内容 + content: String, +} + +#[pymethods] +impl ChatMessage { + #[new] + #[pyo3(signature = (role, content))] + fn new(role: String, content: String) -> Self { + Self { role, content } + } + + #[getter] + fn role(&self) -> String { + self.role.clone() + } + + #[getter] + fn content(&self) -> String { + self.content.clone() + } + + fn __repr__(&self) -> String { + format!("ChatMessage(role={:?}, content={:?})", self.role, self.content) + } +} + +impl ChatMessage { + /// 将 Python 侧消息(dict 或 ChatMessage)转换为 core `ChatMessage` + /// + /// 支持 dict 形式 `{"role": "user", "content": "..."}` 与 `ChatMessage` 实例。 + /// 仅处理文本内容;多模态内容后续阶段支持。 + fn to_core(obj: &Bound<'_, PyAny>) -> PyResult { + if let Ok(msg) = obj.extract::>() { + return Self::from_role_content(&msg.role, &msg.content); + } + let dict: std::collections::HashMap = obj + .extract() + .map_err(|_| { + pyo3::exceptions::PyTypeError::new_err( + "消息必须是 ChatMessage 或含 role/content 的 dict", + ) + })?; + let role = dict + .get("role") + .ok_or_else(|| pyo3::exceptions::PyTypeError::new_err("消息缺少 role 字段"))?; + let content = dict + .get("content") + .ok_or_else(|| pyo3::exceptions::PyTypeError::new_err("消息缺少 content 字段"))?; + Self::from_role_content(role, content) + } + + /// 按角色构造 core `ChatMessage`(仅文本) + fn from_role_content(role: &str, content: &str) -> PyResult { + match role { + "system" => Ok(CoreChatMessage::system(content)), + "user" => Ok(CoreChatMessage::user(content)), + "assistant" => Ok(CoreChatMessage::assistant(content)), + other => Err(pyo3::exceptions::PyTypeError::new_err(format!( + "不支持的消息角色: {other}(阶段 0.6 仅支持 system/user/assistant)" + ))), + } + } +} + +/// 对话完成结果 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ChatCompletion { + id: String, + model: String, + choices: Vec, +} + +#[pymethods] +impl ChatCompletion { + #[getter] + fn id(&self) -> String { + self.id.clone() + } + + #[getter] + fn model(&self) -> String { + self.model.clone() + } + + #[getter] + fn choices(&self) -> Vec { + self.choices.clone() + } + + fn __repr__(&self) -> String { + format!("ChatCompletion(id={:?}, model={:?})", self.id, self.model) + } +} + +impl ChatCompletion { + /// 从 core `ChatCompletion` 构造 Python 包装 + fn from_core(c: CoreChatCompletion) -> Self { + Self { + id: c.id, + model: c.model, + choices: c + .choices + .into_iter() + .map(|ch| ChatChoice { + index: ch.index, + message: ChoiceMessage { + role: ch.message.role, + content: ch.message.content.unwrap_or_default(), + finish_reason: ch.finish_reason.unwrap_or_default(), + }, + }) + .collect(), + } + } +} + +/// 对话选项(choices 元素) +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ChatChoice { + index: u32, + message: ChoiceMessage, +} + +#[pymethods] +impl ChatChoice { + #[getter] + fn index(&self) -> u32 { + self.index + } + + #[getter] + fn message(&self) -> ChoiceMessage { + self.message.clone() + } + + fn __repr__(&self) -> String { + format!("ChatChoice(index={})", self.index) + } +} + +/// 选项中的消息 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ChoiceMessage { + role: String, + content: String, + finish_reason: String, +} + +#[pymethods] +impl ChoiceMessage { + #[getter] + fn role(&self) -> String { + self.role.clone() + } + + #[getter] + fn content(&self) -> String { + self.content.clone() + } + + #[getter] + fn finish_reason(&self) -> String { + self.finish_reason.clone() + } + + fn __repr__(&self) -> String { + format!("ChoiceMessage(role={:?}, content={:?})", self.role, self.content) + } +} + +/// 流式对话块 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ChatCompletionChunk { + id: String, + model: String, + choices: Vec, +} + +#[pymethods] +impl ChatCompletionChunk { + #[getter] + fn id(&self) -> String { + self.id.clone() + } + + #[getter] + fn model(&self) -> String { + self.model.clone() + } + + #[getter] + fn choices(&self) -> Vec { + self.choices.clone() + } + + fn __repr__(&self) -> String { + format!("ChatCompletionChunk(id={:?}, model={:?})", self.id, self.model) + } +} + +impl ChatCompletionChunk { + fn from_core(c: CoreChatCompletionChunk) -> Self { + Self { + id: c.id, + model: c.model, + choices: c + .choices + .into_iter() + .map(|d| ChatChunkDelta { + index: d.index, + role: d.delta.role.unwrap_or_default(), + content: d.delta.content.unwrap_or_default(), + finish_reason: d.finish_reason.unwrap_or_default(), + }) + .collect(), + } + } +} + +/// 流式块中的增量 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ChatChunkDelta { + index: u32, + role: String, + content: String, + finish_reason: String, +} + +#[pymethods] +impl ChatChunkDelta { + #[getter] + fn index(&self) -> u32 { + self.index + } -/// Python 模块入口:import aibridge + #[getter] + fn role(&self) -> String { + self.role.clone() + } + + #[getter] + fn content(&self) -> String { + self.content.clone() + } + + #[getter] + fn finish_reason(&self) -> String { + self.finish_reason.clone() + } + + fn __repr__(&self) -> String { + format!("ChatChunkDelta(index={}, content={:?})", self.index, self.content) + } +} + +/// 文字转语音结果 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct SpeechResult { + /// 音频二进制数据(bytes) + audio_data: Vec, + /// 音频 MIME 类型 + content_type: String, + /// 音频格式(mp3/wav 等) + format: String, + /// 估计音频时长(秒) + duration: Option, + /// 使用的模型 ID + model: Option, +} + +#[pymethods] +impl SpeechResult { + /// 音频二进制数据(Python `bytes`) + #[getter] + fn audio_data<'py>(&self, py: Python<'py>) -> Bound<'py, PyBytes> { + PyBytes::new(py, &self.audio_data) + } + + /// 音频数据长度(便捷访问) + #[getter] + fn size(&self) -> usize { + self.audio_data.len() + } + + #[getter] + fn content_type(&self) -> String { + self.content_type.clone() + } + + #[getter] + fn format(&self) -> String { + self.format.clone() + } + + #[getter] + fn duration(&self) -> Option { + self.duration + } + + #[getter] + fn model(&self) -> Option { + self.model.clone() + } + + fn __repr__(&self) -> String { + format!("SpeechResult(size={}, format={:?})", self.audio_data.len(), self.format) + } +} + +impl SpeechResult { + fn from_core(r: CoreSpeechResult) -> Self { + Self { + audio_data: r.audio_data.unwrap_or_default(), + content_type: r.content_type, + format: r.format, + duration: r.duration, + model: r.model, + } + } +} + +// =========================================================================== +// 流式迭代器 +// =========================================================================== + +/// 流式对话迭代器 +/// +/// 由 `Client.chat_stream` 返回,实现 `__aiter__`/`__anext__` 协议。 +/// `async for chunk in stream:` 每次取一个 `ChatCompletionChunk`,流结束抛 +/// `StopAsyncIteration`。 +/// +/// 实现说明(PyO3 0.28 限制): +/// `__anext__` slot 不支持 `async fn`,故采用"同步 `__anext__` 返回 awaitable" +/// 模式。`__anext__` 在全局 tokio runtime 上 `block_on` 取下一个 chunk(echo +/// adapter 纯计算瞬时完成;真实 adapter 阶段将改为 asyncio.Future 桥接避免 +/// 阻塞事件循环),把结果封装进 [`NextAwaitable`] 返回。 +#[pyclass] +struct ChatStreamIterator { + /// core 流(None 表示已耗尽) + inner: Arc>>, +} + +#[pymethods] +impl ChatStreamIterator { + /// 返回自身(异步迭代器协议:`__aiter__` 返回 self) + fn __aiter__(slf: Py) -> Py { + slf + } + + /// 取下一个 chunk(同步返回 awaitable) + /// + /// 返回一个 [`NextAwaitable`],`await` 后得到 `ChatCompletionChunk` 或抛 + /// `StopAsyncIteration`(流结束)。 + fn __anext__(&self, py: Python<'_>) -> PyResult> { + let inner = self.inner.clone(); + // 在全局 tokio runtime 上同步取下一个 chunk。 + // echo adapter 无 IO,瞬时完成;block_on 在当前线程仅等待 JoinHandle, + // 实际 stream.next() 在 tokio worker 线程执行,不会死锁。 + // py.detach 释放 GIL 期间阻塞,避免长时间持锁。 + let item: Option> = + py.detach(|| RUNTIME.block_on(async move { + let mut guard = inner.lock().await; + match guard.as_mut() { + None => None, + Some(stream) => stream.next().await, + } + })); + + let chunk = match item { + None => None, + Some(Ok(c)) => Some(Ok(Py::new(py, ChatCompletionChunk::from_core(c))?)), + Some(Err(e)) => Some(Err(map_error(e))), + }; + + Py::new(py, NextAwaitable { chunk }) + } +} + +/// `__anext__` 返回的 awaitable +/// +/// 实现 `__await__`/`__iter__`/`__next__` 协议:`await` 时 `__next__` 抛 +/// `StopIteration(chunk)` 返回 chunk,或 `StopAsyncIteration` 表示流结束。 +/// +/// 结果在构造时预计算(由 `__anext__` 的 block_on 完成)。 +#[pyclass] +struct NextAwaitable { + /// 预计算的下一个 chunk(None 表示流已结束) + chunk: Option>>, +} + +#[pymethods] +impl NextAwaitable { + /// `__await__` 返回 self(awaitable 协议) + fn __await__(slf: Py) -> Py { + slf + } + + /// `__iter__` 返回 self(awaitable 兼容迭代器协议) + fn __iter__(slf: Py) -> Py { + slf + } + + /// `__next__` 抛出结果 + /// + /// - 有 chunk:抛 `StopIteration(chunk)`,`await` 得到 chunk + /// - 流结束:抛 `StopAsyncIteration` + /// - 取 chunk 出错:抛对应 AibridgeError 子类 + fn __next__(&self, py: Python<'_>) -> PyResult> { + match &self.chunk { + None => Err(pyo3::exceptions::PyStopAsyncIteration::new_err(())), + Some(Ok(chunk)) => { + // StopIteration(chunk) → await 得到 chunk + Err(pyo3::exceptions::PyStopIteration::new_err((chunk.clone_ref(py),))) + } + Some(Err(e)) => Err(e.clone_ref(py)), + } + } +} + +// =========================================================================== +// 客户端 +// =========================================================================== + +/// AIBridge 统一客户端 +/// +/// 对应 Python v1 `Client`,是用户使用 SDK 的唯一入口。 +/// +/// 示例: +/// ```python +/// import asyncio +/// from aibridge import Client +/// +/// async def main(): +/// client = Client(provider="echo") +/// await client.start() +/// resp = await client.chat(model="echo-chat", +/// messages=[{"role": "user", "content": "hello"}]) +/// print(resp.choices[0].message.content) +/// await client.close() +/// +/// asyncio.run(main()) +/// ``` +#[pyclass] +struct Client { + /// core 客户端(用 tokio Mutex 保护以支持 start/close 可变操作) + inner: Arc>, + /// Provider 类型(构造后不变,缓存以避免同步 getter 中 block_on) + provider_type: String, +} + +#[pymethods] +impl Client { + /// 创建客户端 + /// + /// 参数: + /// - `provider`: Provider 类型(如 "echo"、"openai") + /// - `api_key`: 可选 API Key(免认证 provider 可省略) + /// - `base_url`: 可选 API Base URL + #[new] + #[pyo3(signature = (provider, *, api_key=None, base_url=None))] + fn new(provider: &str, api_key: Option, base_url: Option) -> PyResult { + let mut opts_builder = CoreClientOptions::builder(); + if let Some(key) = api_key { + opts_builder = opts_builder.api_key(key); + } + if let Some(url) = base_url { + opts_builder = opts_builder.base_url(url); + } + let opts = opts_builder.build(); + let core_client = CoreClient::new(provider, opts).map_err(map_error)?; + let provider_type = core_client.provider_type().to_string(); + Ok(Self { + inner: Arc::new(Mutex::new(core_client)), + provider_type, + }) + } + + /// Provider 类型 + #[getter] + fn provider_type(&self) -> String { + self.provider_type.clone() + } + + /// 启动客户端(初始化适配器) + async fn start(&self) -> PyResult<()> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let mut client = inner.lock().await; + client.start().await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("start 任务失败: {e}")) + })?; + result.map_err(map_error) + } + + /// 关闭客户端(释放资源) + async fn close(&self) -> PyResult<()> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let mut client = inner.lock().await; + client.close().await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("close 任务失败: {e}")) + })?; + result.map_err(map_error) + } + + /// 文本对话 + /// + /// 参数: + /// - `model`: 模型名称 + /// - `messages`: 消息列表(`ChatMessage` 或 `{"role":..., "content":...}` dict) + /// - `temperature`: 可选温度系数 + /// - `max_tokens`: 可选最大 token 数 + #[pyo3(signature = (model, messages, *, temperature=None, max_tokens=None))] + async fn chat( + &self, + model: String, + messages: Vec>, + temperature: Option, + max_tokens: Option, + ) -> PyResult { + // 在持有 GIL 时把 Python 消息转换为 core 消息 + let core_messages = Python::attach(|py| { + messages + .iter() + .map(|m| ChatMessage::to_core(m.bind(py))) + .collect::>>() + })?; + + let mut builder = CoreChatRequest::builder(model, core_messages); + if let Some(t) = temperature { + builder = builder.temperature(t); + } + if let Some(m) = max_tokens { + builder = builder.max_tokens(m); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.chat(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("chat 任务失败: {e}")) + })?; + + let completion = result.map_err(map_error)?; + Ok(ChatCompletion::from_core(completion)) + } + + /// 流式文本对话 + /// + /// 返回 `ChatStreamIterator`(异步迭代器),`async for chunk in ...` 逐块消费。 + /// + /// 参数同 `chat`。 + #[pyo3(signature = (model, messages, *, temperature=None, max_tokens=None))] + async fn chat_stream( + &self, + model: String, + messages: Vec>, + temperature: Option, + max_tokens: Option, + ) -> PyResult { + let core_messages = Python::attach(|py| { + messages + .iter() + .map(|m| ChatMessage::to_core(m.bind(py))) + .collect::>>() + })?; + + let mut builder = CoreChatRequest::builder(model, core_messages).stream(true); + if let Some(t) = temperature { + builder = builder.temperature(t); + } + if let Some(m) = max_tokens { + builder = builder.max_tokens(m); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let stream_result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.chat_stream(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!( + "chat_stream 任务失败: {e}" + )) + })?; + + let stream = stream_result.map_err(map_error)?; + Ok(ChatStreamIterator { + inner: Arc::new(Mutex::new(Some(stream))), + }) + } + + /// 文字转语音 + /// + /// 参数: + /// - `model`: TTS 模型名称 + /// - `input`: 要合成的文本 + /// - `voice`: 音色(字符串) + /// - `response_format`: 可选音频格式(默认 mp3) + /// - `speed`: 可选语速(0.25-4.0) + #[pyo3(signature = (model, input, voice, *, response_format=None, speed=None))] + async fn speech( + &self, + model: String, + input: String, + voice: String, + response_format: Option, + speed: Option, + ) -> PyResult { + let mut builder = CoreSpeechRequest::builder(model, input, voice); + if let Some(f) = response_format { + builder = builder.response_format(f); + } + if let Some(s) = speed { + builder = builder.speed(s); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.speech(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("speech 任务失败: {e}")) + })?; + + let speech = result.map_err(map_error)?; + Ok(SpeechResult::from_core(speech)) + } +} + +// =========================================================================== +// 模块入口 +// =========================================================================== + +/// Python 模块入口:`import aibridge` #[pymodule] -fn aibridge(_py: Python, _m: &Bound) -> PyResult<()> { - // 阶段 0.6 填充:注册 Client / Router / 错误类 / 数据模型 +fn aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { + let py = m.py(); + + // 触发全局 runtime 初始化(首次访问 Lazy 即建) + let _ = &*RUNTIME; + + // 错误类(基类 + 子类) + m.add("AibridgeError", py.get_type::())?; + m.add("AuthenticationError", py.get_type::())?; + m.add("RateLimitError", py.get_type::())?; + m.add("ValidationError", py.get_type::())?; + m.add("ModelNotFoundError", py.get_type::())?; + m.add("APIError", py.get_type::())?; + m.add("NetworkError", py.get_type::())?; + m.add("TimeoutError", py.get_type::())?; + m.add( + "UnsupportedCapabilityError", + py.get_type::(), + )?; + m.add("ProviderNotFoundError", py.get_type::())?; + m.add("VoiceNotAvailableError", py.get_type::())?; + m.add( + "ServiceUnavailableError", + py.get_type::(), + )?; + + // 数据模型 + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + + // 客户端与流式 + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + + // 模块版本 + m.add("__version__", aibridge_core::VERSION)?; + Ok(()) } diff --git a/examples/hello_python.py b/examples/hello_python.py new file mode 100644 index 0000000..4a4e17c --- /dev/null +++ b/examples/hello_python.py @@ -0,0 +1,85 @@ +"""AIBridge Python 绑定 hello world + +验证 PyO3 直连 aibridge-core 的端到端管线: +- chat:回显最后 user 消息 + " [echo]"(期望 "hello [echo]") +- chat_stream:流式产 3 个 chunk +- speech:返回 15 字节固定音频 + +使用 echo mock 适配器,免认证,无网络依赖。 +""" + +import asyncio + +import aibridge +from aibridge import Client + + +async def main() -> None: + print(f"aibridge 版本: {aibridge.__version__}") + print("=" * 50) + + # 创建 echo 客户端(免认证) + client = Client(provider="echo") + await client.start() + print(f"已创建客户端,provider_type={client.provider_type}") + + # --- chat --- + print("-" * 50) + print("[chat] 调用 chat(model='echo-chat', messages=[{'role':'user','content':'hello'}])") + resp = await client.chat( + model="echo-chat", + messages=[{"role": "user", "content": "hello"}], + ) + content = resp.choices[0].message.content + print(f"[chat] choices[0].message.content = {content!r}") + assert content == "hello [echo]", f"期望 'hello [echo]',实际 {content!r}" + print("[chat] 通过:回显正确") + + # --- chat_stream --- + print("-" * 50) + print("[stream] 调用 chat_stream(...),逐块消费") + stream = await client.chat_stream( + model="echo-chat", + messages=[{"role": "user", "content": "hello"}], + ) + chunks = [] + async for chunk in stream: + delta = chunk.choices[0] + print( + f"[stream] chunk: index={delta.index} role={delta.role!r} " + f"content={delta.content!r} finish_reason={delta.finish_reason!r}" + ) + chunks.append(delta) + assert len(chunks) == 3, f"期望 3 个 chunk,实际 {len(chunks)}" + # 拼接第 2、3 块内容应等于完整回显 + assembled = chunks[1].content + chunks[2].content + assert assembled == "hello [echo]", f"拼接内容 {assembled!r} 不等于 'hello [echo]'" + assert chunks[0].role == "assistant", "第 1 块应含 role='assistant'" + assert chunks[2].finish_reason == "stop", "第 3 块应标记 finish_reason='stop'" + print(f"[stream] 通过:3 个 chunk,拼接 = {assembled!r}") + + # --- speech --- + print("-" * 50) + print("[speech] 调用 speech(model='echo-tts', input='hello', voice='alloy')") + result = await client.speech( + model="echo-tts", + input="hello", + voice="alloy", + ) + audio = result.audio_data + print(f"[speech] len(audio_data) = {len(audio)}") + print(f"[speech] content_type={result.content_type} format={result.format}") + assert len(audio) == 15, f"期望 15 字节,实际 {len(audio)}" + print(f"[speech] 通过:音频 {len(audio)} 字节") + + # --- 关闭 --- + print("-" * 50) + await client.close() + print("[close] 客户端已关闭") + + print("=" * 50) + print("全部通过:chat 回显 + stream 3 chunk + speech 15 字节") + + +if __name__ == "__main__": + asyncio.run(main()) From 0a802b96193c14211457d954561e4396bce9cd33 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 12:13:06 +0800 Subject: [PATCH 11/55] =?UTF-8?q?feat(aibridge-dotnet):=20=E9=98=B6?= =?UTF-8?q?=E6=AE=B50.6=20P/Invoke=20=E7=BB=91=E5=AE=9A?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 实现 .NET (C#) 绑定,通过 P/Invoke 调 aibridge-ffi cdylib。 实现文件: - AIBridge.csproj:net8.0 类库 + 可执行 hello(OutputType=Exe, StartupObject=Hello) - Native.cs:P/Invoke 声明全部 aibridge_* 函数 + SafeHandle 封装(字符串/字节自动释放)+ 运行时动态库解析(支持 AIBRIDGE_LIB_PATH 环境变量与仓库 target/debug 回退) - Models.cs:数据模型(ChatRequest/ChatMessage/ChatCompletion/ChatCompletionChunk/SpeechRequest/SpeechResult 等),System.Text.Json 反序列化,字段名对齐 Rust serde snake_case - AibridgeException.cs:异常体系(AibridgeException + 11 个子类,按 last_error 的 code 字段映射) - Client.cs:Client 封装(IDisposable),Chat/Speech/ChatStreamAsync(IAsyncEnumerable),async 方法用 Task.Run 包装阻塞 FFI - Hello.cs:hello world 示例(Program Main),验证 echo 适配器 chat 回显 + stream 3 chunk + speech 15 字节 - build.sh:构建脚本(cargo build ffi + dotnet build + dotnet run) FFI 遗留问题处理: 1. last_error 线程局部:每个 FFI 失败后同线程立即读 last_error 转存(ChatStreamAsync 在 Task.Run 同一 ThreadPool 线程内读取) 2. stream_next 串行:await foreach 天然串行,不并发 next 3. aibridge_bytes_free / aibridge_string_free:用 AibridgeBytesHandle / AibridgeStringHandle(SafeHandle)兜底释放,MarshalAndFree 先拷贝后释放 4. client / stream 句柄:Client.Dispose(Interlocked 原子化防 double-free)调 aibridge_client_destroy;ChatStreamAsync finally 调 aibridge_stream_destroy 验收: - cargo build -p aibridge-ffi 已产 libaibridge.dylib(target/debug/) - dotnet SDK 未装(沙箱无法 sudo 安装),hello world 待 dotnet 环境验证 - build.sh 提供完整构建+运行流程,装好 dotnet 后执行 ./build.sh run 即可 --- bindings/dotnet/.gitignore | 13 + bindings/dotnet/AIBridge/AIBridge.csproj | 53 +++ bindings/dotnet/AIBridge/AibridgeException.cs | 186 ++++++++++ bindings/dotnet/AIBridge/Client.cs | 316 +++++++++++++++++ bindings/dotnet/AIBridge/Hello.cs | 94 +++++ bindings/dotnet/AIBridge/Models.cs | 277 +++++++++++++++ bindings/dotnet/AIBridge/Native.cs | 320 ++++++++++++++++++ bindings/dotnet/build.sh | 112 ++++++ 8 files changed, 1371 insertions(+) create mode 100644 bindings/dotnet/.gitignore create mode 100644 bindings/dotnet/AIBridge/AIBridge.csproj create mode 100644 bindings/dotnet/AIBridge/AibridgeException.cs create mode 100644 bindings/dotnet/AIBridge/Client.cs create mode 100644 bindings/dotnet/AIBridge/Hello.cs create mode 100644 bindings/dotnet/AIBridge/Models.cs create mode 100644 bindings/dotnet/AIBridge/Native.cs create mode 100755 bindings/dotnet/build.sh diff --git a/bindings/dotnet/.gitignore b/bindings/dotnet/.gitignore new file mode 100644 index 0000000..2c4763c --- /dev/null +++ b/bindings/dotnet/.gitignore @@ -0,0 +1,13 @@ +# .NET 构建产物 +bin/ +obj/ + +# IDE +.vs/ +.vscode/ +.idea/ +*.user +*.suo + +# macOS +.DS_Store diff --git a/bindings/dotnet/AIBridge/AIBridge.csproj b/bindings/dotnet/AIBridge/AIBridge.csproj new file mode 100644 index 0000000..9f92aa1 --- /dev/null +++ b/bindings/dotnet/AIBridge/AIBridge.csproj @@ -0,0 +1,53 @@ + + + + + + net8.0 + latest + enable + enable + AIBridge + AIBridge + + + Exe + AIBridge.Hello + + + $(MSBuildThisFileDirectory)../../../target/debug + + + + + + + + + + diff --git a/bindings/dotnet/AIBridge/AibridgeException.cs b/bindings/dotnet/AIBridge/AibridgeException.cs new file mode 100644 index 0000000..e92d70f --- /dev/null +++ b/bindings/dotnet/AIBridge/AibridgeException.cs @@ -0,0 +1,186 @@ +using System.Text.Json; + +namespace AIBridge; + +// ============================================================================ +// 异常体系 +// +// 对应设计文档第 9 节 .NET 异常映射:AibridgeException + 子类。 +// 子类与 core AibridgeError 枚举变体一一对应(rate_limit/authentication/...)。 +// +// FFI 错误模型(设计文档 7.4):aibridge_status_t 返回码 + aibridge_last_error() +// 线程局部 JSON:{"code":"...","message":"...","details":...,"retryable":bool}。 +// Client 在 FFI 失败后同线程读取 last_error 转存,按 code 映射为子类异常。 +// ============================================================================ + +///

AIBridge 异常基类。 +public class AibridgeException : Exception +{ + /// 错误码字符串(如 "rate_limit"、"authentication")。 + public string Code { get; } + + /// 是否可重试。 + public bool Retryable { get; } + + /// 原始 details(可能为 null)。 + public JsonElement? Details { get; } + + public AibridgeException(string message, string code = "unknown", bool retryable = false, + JsonElement? details = null, Exception? inner = null) + : base(message, inner) + { + Code = code; + Retryable = retryable; + Details = details; + } + + /// 内部构造:仅状态码已知、无 last_error JSON 时,按状态码映射默认子类。 + internal static AibridgeException FromStatus(int status, string? lastErrorJson) + { + // 解析 last_error JSON(同线程立即读,符合 FFI 线程局部语义) + string code = "ffi_error"; + string message = $"FFI 调用失败,状态码={status}"; + bool retryable = false; + JsonElement? details = null; + + if (!string.IsNullOrEmpty(lastErrorJson)) + { + try + { + using JsonDocument doc = JsonDocument.Parse(lastErrorJson); + JsonElement root = doc.RootElement; + if (root.TryGetProperty("code", out JsonElement c)) code = c.GetString() ?? code; + if (root.TryGetProperty("message", out JsonElement m)) + message = m.GetString() ?? message; + if (root.TryGetProperty("retryable", out JsonElement r)) retryable = r.GetBoolean(); + if (root.TryGetProperty("details", out JsonElement d)) details = d.Clone(); + } + catch (JsonException) + { + // last_error 不是合法 JSON,回退用状态码 + 原始字符串 + message = lastErrorJson; + } + } + + // 先按 code(来自 last_error)映射,更精确 + AibridgeException ex = MapByCode(code, message, retryable, details); + if (ex != null) return ex; + + // code 未识别时按状态码兜底 + return status switch + { + AibridgeStatus.Authentication => new AuthenticationException(message, retryable, details), + AibridgeStatus.RateLimit => new RateLimitException(message, retryable, details), + AibridgeStatus.Validation => new ValidationException(message, retryable, details), + AibridgeStatus.ModelNotFound => new ModelNotFoundException(message, retryable, details), + AibridgeStatus.Api => new ApiException(message, retryable, details), + AibridgeStatus.Network => new NetworkException(message, retryable, details), + AibridgeStatus.Timeout => new TimeoutException_(message, retryable, details), + AibridgeStatus.UnsupportedCapability => new UnsupportedCapabilityException(message, retryable, details), + AibridgeStatus.ProviderNotFound => new ProviderNotFoundException(message, retryable, details), + AibridgeStatus.VoiceNotAvailable => new VoiceNotAvailableException(message, retryable, details), + AibridgeStatus.ServiceUnavailable => new ServiceUnavailableException(message, retryable, details), + _ => new AibridgeException(message, code, retryable, details), + }; + } + + /// 按 last_error 的 code 字段映射子类(core AibridgeError 变体名)。 + private static AibridgeException? MapByCode(string code, string message, bool retryable, JsonElement? details) + { + return code switch + { + "authentication" => new AuthenticationException(message, retryable, details), + "rate_limit" => new RateLimitException(message, retryable, details), + "validation_error" => new ValidationException(message, retryable, details), + "model_not_found" => new ModelNotFoundException(message, retryable, details), + "api_error" => new ApiException(message, retryable, details), + "network_error" => new NetworkException(message, retryable, details), + "timeout" => new TimeoutException_(message, retryable, details), + "unsupported_capability" => new UnsupportedCapabilityException(message, retryable, details), + "provider_not_found" => new ProviderNotFoundException(message, retryable, details), + "voice_not_available" => new VoiceNotAvailableException(message, retryable, details), + "service_unavailable" => new ServiceUnavailableException(message, retryable, details), + _ => null, // 未识别 code,交由状态码兜底 + }; + } +} + +// —— 子类(与 core AibridgeError 变体一一对应)—————————————— + +public class AuthenticationException : AibridgeException +{ + public AuthenticationException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "authentication", retryable, details) { } +} + +public class RateLimitException : AibridgeException +{ + /// 建议等待秒数(若 provider 返回)。 + public double? RetryAfter { get; } + + public RateLimitException(string msg, bool retryable = true, JsonElement? details = null) + : base(msg, "rate_limit", retryable, details) + { + // 尝试从 details.retry_after 读取 + if (details.HasValue && details.Value.TryGetProperty("retry_after", out JsonElement r) + && r.ValueKind == JsonValueKind.Number) + { + RetryAfter = r.GetDouble(); + } + } +} + +public class ValidationException : AibridgeException +{ + public ValidationException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "validation_error", retryable, details) { } +} + +public class ModelNotFoundException : AibridgeException +{ + public ModelNotFoundException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "model_not_found", retryable, details) { } +} + +public class ApiException : AibridgeException +{ + public ApiException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "api_error", retryable, details) { } +} + +public class NetworkException : AibridgeException +{ + public NetworkException(string msg, bool retryable = true, JsonElement? details = null) + : base(msg, "network_error", retryable, details) { } +} + +/// 避免与 System.TimeoutException 重名,加下划线后缀。 +public class TimeoutException_ : AibridgeException +{ + public TimeoutException_(string msg, bool retryable = true, JsonElement? details = null) + : base(msg, "timeout", retryable, details) { } +} + +public class UnsupportedCapabilityException : AibridgeException +{ + public UnsupportedCapabilityException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "unsupported_capability", retryable, details) { } +} + +public class ProviderNotFoundException : AibridgeException +{ + public ProviderNotFoundException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "provider_not_found", retryable, details) { } +} + +public class VoiceNotAvailableException : AibridgeException +{ + public VoiceNotAvailableException(string msg, bool retryable = false, JsonElement? details = null) + : base(msg, "voice_not_available", retryable, details) { } +} + +public class ServiceUnavailableException : AibridgeException +{ + public ServiceUnavailableException(string msg, bool retryable = true, JsonElement? details = null) + : base(msg, "service_unavailable", retryable, details) { } +} diff --git a/bindings/dotnet/AIBridge/Client.cs b/bindings/dotnet/AIBridge/Client.cs new file mode 100644 index 0000000..10aaf62 --- /dev/null +++ b/bindings/dotnet/AIBridge/Client.cs @@ -0,0 +1,316 @@ +using System.Runtime.InteropServices; +using System.Text; +using System.Text.Json; + +namespace AIBridge; + +// ============================================================================ +// Client 封装层 +// +// 包装 aibridge_client_t* 句柄,提供 Chat / ChatStreamAsync / Speech 三个能力。 +// +// 设计要点(FFI 遗留问题全部处理): +// 1. last_error 线程局部:每次 FFI 失败后,在【同一托管线程】立即调 +// aibridge_last_error() 读取并转存为字符串,再交给 AibridgeException.FromStatus。 +// 注意:因 ReadLastError 与 FFI 调用必须同线程,这里用普通同步方法实现核心逻辑, +// 再用 Task.Run 包装为异步暴露给上层(保证 FFI 调用 + last_error 读取在同一 +// ThreadPool 线程,不跨线程)。 +// 2. stream_next 串行:ChatStreamAsync 的迭代器逐个 await,天然串行,不并发 next。 +// 3. 字符串/字节释放:用 AibridgeStringHandle / AibridgeBytesHandle(SafeHandle)兜底。 +// 4. 句柄释放:Client 实现 IDisposable,调 aibridge_client_destroy;Stream 同理。 +// ============================================================================ + +/// +/// AIBridge 客户端。对应 core Client,通过 P/Invoke 调 aibridge-ffi。 +/// +public sealed class Client : IDisposable +{ + private IntPtr _handle; // aibridge_client_t* + private bool _started; + private int _disposed; // 0=未释放,1=已释放(用 Interlocked 原子操作) + + // JSON 序列化选项:与 Rust serde 默认行为对齐。 + // 字段名已用 JsonPropertyName 显式指定 snake_case,不依赖命名策略转换。 + private static readonly JsonSerializerOptions JsonOpts = new() + { + DefaultIgnoreCondition = System.Text.Json.Serialization.JsonIgnoreCondition.WhenWritingNull, + PropertyNamingPolicy = null, + // 容错:允许尾随逗号与注释(provider 响应可能含) + ReadCommentHandling = JsonCommentHandling.Skip, + AllowTrailingCommas = true, + }; + + /// 创建客户端(对应 aibridge_client_new)。 + /// Provider 类型(如 "echo"、"openai"、"agnes")。 + /// ClientOptions 的 JSON,可为 null(默认配置)。 + public Client(string provider, string? configJson = null) + { + if (string.IsNullOrEmpty(provider)) + throw new ArgumentException("provider 不能为空", nameof(provider)); + + // C# string → UTF-8 字节数组(含 NUL 终止),匹配 C 字符串契约 + byte[] providerBytes = ToCString(provider); + byte[]? configBytes = configJson != null ? ToCString(configJson) : null; + + IntPtr handle = Native.aibridge_client_new(providerBytes, configBytes); + if (handle == IntPtr.Zero) + { + // 失败:同线程立即读 last_error(线程局部,不可跨线程) + throw AibridgeException.FromStatus(AibridgeStatus.Ffi, ReadLastError()); + } + _handle = handle; + } + + /// 启动客户端(对应 aibridge_client_start)。 + public void Start() + { + ThrowIfDisposed(); + if (_started) return; + + int status = Native.aibridge_client_start(_handle); + if (status != AibridgeStatus.Ok) + { + throw AibridgeException.FromStatus(status, ReadLastError()); + } + _started = true; + } + + /// 启动客户端(异步包装)。 + public Task StartAsync(CancellationToken cancellationToken = default) + => Task.Run(Start, cancellationToken); + + // —— 文本对话 —————————————————————————————————————————— + + /// 文本对话(阻塞,对应 aibridge_client_chat)。 + public ChatCompletion Chat(ChatRequest request) + { + ThrowIfDisposed(); + ArgumentNullException.ThrowIfNull(request); + + byte[] reqJson = ToCString(JsonSerializer.Serialize(request, JsonOpts)); + IntPtr outResponse = IntPtr.Zero; + + int status = Native.aibridge_client_chat(_handle, reqJson, ref outResponse); + // 同线程立即读 last_error(线程局部,必须与 FFI 调用同线程) + string? lastError = ReadLastError(); + + if (status != AibridgeStatus.Ok) + { + // 失败时 outResponse 应为 Zero,但防御性释放 + if (outResponse != IntPtr.Zero) Native.aibridge_string_free(outResponse); + throw AibridgeException.FromStatus(status, lastError); + } + + // 成功:用 SafeHandle 接管字符串,拷贝后释放 + var handle = new AibridgeStringHandle(); + handle.SetHandle(outResponse); + string? responseJson = handle.MarshalAndFree(); + + if (string.IsNullOrEmpty(responseJson)) + { + throw new AibridgeException("chat 返回空响应"); + } + + return JsonSerializer.Deserialize(responseJson, JsonOpts) + ?? throw new AibridgeException("反序列化 ChatCompletion 失败"); + } + + /// 文本对话(异步包装,用 Task.Run 在 ThreadPool 调度阻塞 FFI)。 + public Task ChatAsync(ChatRequest request, CancellationToken cancellationToken = default) + => Task.Run(() => Chat(request), cancellationToken); + + // —— 流式对话 —————————————————————————————————————————— + + /// + /// 流式文本对话(对应 aibridge_client_chat_stream + stream_next 循环)。 + /// + /// 返回 IAsyncEnumerable,可用 await foreach 消费。 + /// 内部 stream_next 串行调用(同一 stream 不可并发 next,符合 FFI 约束)。 + /// 取消通过 CancellationToken:迭代器在下次 next 前检查,并提前 destroy stream。 + /// + public async IAsyncEnumerable ChatStreamAsync( + ChatRequest request, + [System.Runtime.CompilerServices.EnumeratorCancellation] CancellationToken cancellationToken = default) + { + ThrowIfDisposed(); + ArgumentNullException.ThrowIfNull(request); + + byte[] reqJson = ToCString(JsonSerializer.Serialize(request, JsonOpts)); + IntPtr outStream = IntPtr.Zero; + + int status; + string? lastError; + // 创建 stream 必须同步完成(FFI 阻塞 + last_error 线程局部) + status = Native.aibridge_client_chat_stream(_handle, reqJson, ref outStream); + lastError = ReadLastError(); + + if (status != AibridgeStatus.Ok) + { + throw AibridgeException.FromStatus(status, lastError); + } + + // 用 try/finally 保证 stream 句柄一定被 destroy(即使中途取消或异常) + try + { + while (true) + { + cancellationToken.ThrowIfCancellationRequested(); + + // stream_next 阻塞,包到 Task.Run 避免阻塞调用方线程。 + // 串行:每次 await 完成后才进入下一次 next,绝不并发。 + // 用元组返回结果,避免实例字段带来的并发隐患(同一 Client 可能有多个 stream)。 + (int nextStatus, string? chunkJson, string? lastError) result = await Task.Run(() => + { + IntPtr outChunk = IntPtr.Zero; + int s = Native.aibridge_stream_next(outStream, ref outChunk); + // 同线程读 last_error(线程局部,必须与 FFI 调用同线程) + string? err = s < 0 ? ReadLastError() : null; + string? json = null; + if (s == AibridgeStatus.StreamChunk && outChunk != IntPtr.Zero) + { + // 拷贝 chunk JSON 并释放原生字符串(SafeHandle 兜底释放) + var h = new AibridgeStringHandle(); + h.SetHandle(outChunk); + json = h.MarshalAndFree(); + } + else if (outChunk != IntPtr.Zero) + { + // 防御性释放:FFI 在非 chunk 路径若意外写入 outChunk,避免泄漏 + Native.aibridge_string_free(outChunk); + } + return (s, json, err); + }, cancellationToken).ConfigureAwait(false); + + if (result.nextStatus == AibridgeStatus.StreamEnd) + { + yield break; // 流正常结束 + } + + if (result.nextStatus < 0) + { + throw AibridgeException.FromStatus(result.nextStatus, result.lastError); + } + + // nextStatus == StreamChunk:反序列化 chunk + if (string.IsNullOrEmpty(result.chunkJson)) + { + throw new AibridgeException("stream_next 返回 chunk 但 JSON 为空"); + } + + ChatCompletionChunk? chunk = JsonSerializer.Deserialize(result.chunkJson, JsonOpts); + if (chunk == null) + { + throw new AibridgeException("反序列化 ChatCompletionChunk 失败"); + } + yield return chunk; + } + } + finally + { + // 无论正常结束、取消、异常,都 destroy stream(触发 Rust drop → tokio task abort) + if (outStream != IntPtr.Zero) + { + Native.aibridge_stream_destroy(outStream); + } + } + } + + // —— 文字转语音 —————————————————————————————————————————— + + /// 文字转语音(阻塞,对应 aibridge_client_speech,二进制走 aibridge_bytes_t)。 + public SpeechResult Speech(SpeechRequest request) + { + ThrowIfDisposed(); + ArgumentNullException.ThrowIfNull(request); + + byte[] reqJson = ToCString(JsonSerializer.Serialize(request, JsonOpts)); + IntPtr outAudio = IntPtr.Zero; + IntPtr outMeta = IntPtr.Zero; + + int status = Native.aibridge_client_speech(_handle, reqJson, ref outAudio, ref outMeta); + // 同线程立即读 last_error + string? lastError = ReadLastError(); + + if (status != AibridgeStatus.Ok) + { + // 防御性释放(FFI 失败时应为 Zero) + if (outAudio != IntPtr.Zero) Native.aibridge_bytes_free(outAudio); + if (outMeta != IntPtr.Zero) Native.aibridge_string_free(outMeta); + throw AibridgeException.FromStatus(status, lastError); + } + + // 二进制音频:用 SafeHandle 接管,拷贝后释放 + byte[] audioData = Array.Empty(); + if (outAudio != IntPtr.Zero) + { + var audioHandle = new AibridgeBytesHandle(); + audioHandle.SetHandle(outAudio); + audioData = audioHandle.MarshalAndFree(); + } + + // 元数据 JSON + SpeechResult? result; + if (outMeta != IntPtr.Zero) + { + var metaHandle = new AibridgeStringHandle(); + metaHandle.SetHandle(outMeta); + string? metaJson = metaHandle.MarshalAndFree(); + result = string.IsNullOrEmpty(metaJson) + ? new SpeechResult() + : JsonSerializer.Deserialize(metaJson, JsonOpts); + } + else + { + result = new SpeechResult(); + } + + result ??= new SpeechResult(); + result.AudioData = audioData; + return result; + } + + /// 文字转语音(异步包装)。 + public Task SpeechAsync(SpeechRequest request, CancellationToken cancellationToken = default) + => Task.Run(() => Speech(request), cancellationToken); + + // —— Dispose 模式 —————————————————————————————————————————— + + public void Dispose() + { + // 原子化:防并发 Dispose 导致 double-free aibridge_client_destroy + if (Interlocked.Exchange(ref _disposed, 1) == 1) return; + + // 原子置换 _handle,确保只释放一次 + IntPtr h = Interlocked.Exchange(ref _handle, IntPtr.Zero); + if (h != IntPtr.Zero) + { + Native.aibridge_client_destroy(h); + } + } + + // —— 内部辅助 ———————————————————————————————————————————— + + /// 读取当前线程的 last_error(线程局部,必须与 FFI 调用同线程)。 + private static string? ReadLastError() + { + IntPtr ptr = Native.aibridge_last_error(); + if (ptr == IntPtr.Zero) return null; + // 拷贝为托管字符串(指针仅在当前线程下一次 FFI 调用前有效) + return Marshal.PtrToStringUTF8(ptr); + } + + /// C# string → UTF-8 字节数组(含 NUL 终止),匹配 C 字符串契约。 + private static byte[] ToCString(string s) + { + // Encoding.UTF8.GetBytes + 额外 NUL 字节 + byte[] bytes = Encoding.UTF8.GetBytes(s); + byte[] withNull = new byte[bytes.Length + 1]; + Buffer.BlockCopy(bytes, 0, withNull, 0, bytes.Length); + return withNull; + } + + private void ThrowIfDisposed() + { + if (Volatile.Read(ref _disposed) != 0) throw new ObjectDisposedException(nameof(Client)); + } +} diff --git a/bindings/dotnet/AIBridge/Hello.cs b/bindings/dotnet/AIBridge/Hello.cs new file mode 100644 index 0000000..92c2f04 --- /dev/null +++ b/bindings/dotnet/AIBridge/Hello.cs @@ -0,0 +1,94 @@ +namespace AIBridge; + +// ============================================================================ +// Hello world 示例(Program Main) +// +// 验证 .NET 绑定全链路:echo 适配器免认证可用 +// 1. new Client("echo") + Start +// 2. Chat(echo-chat, [user:hello]) → 打印 choices[0].message.content(期望 "hello [echo]") +// 3. ChatStreamAsync 同上 → await foreach 打印 chunk(期望 3 个) +// 4. Speech(echo-tts, input:hello, voice:alloy) → 打印 AudioData.Length(期望 15) +// 5. Dispose +// +// 运行方式:dotnet run(需先 cargo build -p aibridge-ffi 产 libaibridge.dylib) +// ============================================================================ + +public static class Hello +{ + public static async Task Main(string[] args) + { + Console.WriteLine("=== AIBridge .NET Hello World (echo adapter) ===\n"); + + try + { + // 1. 创建 + 启动客户端(echo 免认证) + using var client = new Client("echo"); + client.Start(); + Console.WriteLine("[OK] Client(echo) created and started.\n"); + + // 2. Chat:echo 回显 +" [echo]" + Console.WriteLine("--- Chat ---"); + var chatReq = new ChatRequest("echo-chat", new[] + { + ChatMessage.User("hello"), + }); + ChatCompletion completion = client.Chat(chatReq); + string? content = completion.Choices.Count > 0 + ? completion.Choices[0].Message.Content + : null; + Console.WriteLine($" choices[0].message.content = {content ?? "(空)"}"); + Console.WriteLine($" 期望: \"hello [echo]\",实际: \"{content}\""); + Console.WriteLine($" 匹配: {content == "hello [echo]"}\n"); + + // 3. ChatStream:echo 返 3 个 chunk + Console.WriteLine("--- ChatStream ---"); + var streamReq = new ChatRequest("echo-chat", new[] + { + ChatMessage.User("hello"), + }); + int chunkCount = 0; + var assembled = new System.Text.StringBuilder(); + await foreach (ChatCompletionChunk chunk in client.ChatStreamAsync(streamReq)) + { + chunkCount++; + string? delta = chunk.Choices.Count > 0 + ? chunk.Choices[0].Delta.Content + : null; + Console.WriteLine($" chunk #{chunkCount}: delta.content={delta ?? "(空)"}"); + if (delta != null) assembled.Append(delta); + } + Console.WriteLine($" 共 {chunkCount} 个 chunk,拼接结果: \"{assembled}\""); + Console.WriteLine($" 期望 3 个 chunk,实际: {chunkCount}\n"); + + // 4. Speech:echo 返 15 字节固定音频 + Console.WriteLine("--- Speech ---"); + var speechReq = new SpeechRequest("echo-tts", "hello", "alloy"); + SpeechResult speech = client.Speech(speechReq); + Console.WriteLine($" AudioData.Length = {speech.AudioData.Length}"); + Console.WriteLine($" format = {speech.Format}, model = {speech.Model}"); + Console.WriteLine($" 期望 15 字节,实际: {speech.AudioData.Length}\n"); + + // 5. Dispose(using 自动调) + Console.WriteLine("[OK] Client disposed."); + + // 汇总 + Console.WriteLine("\n=== 汇总 ==="); + Console.WriteLine($" Chat 回显匹配: {content == "hello [echo]"}"); + Console.WriteLine($" Stream chunk 数匹配(==3): {chunkCount == 3}"); + Console.WriteLine($" Speech 字节数匹配(==15): {speech.AudioData.Length == 15}"); + return 0; + } + catch (AibridgeException ex) + { + Console.Error.WriteLine($"[FAIL] AibridgeException: code={ex.Code}, retryable={ex.Retryable}"); + Console.Error.WriteLine($" message: {ex.Message}"); + return 1; + } + catch (Exception ex) + { + Console.Error.WriteLine($"[FAIL] {ex.GetType().Name}: {ex.Message}"); + Console.Error.WriteLine(ex); + return 2; + } + } +} diff --git a/bindings/dotnet/AIBridge/Models.cs b/bindings/dotnet/AIBridge/Models.cs new file mode 100644 index 0000000..7786257 --- /dev/null +++ b/bindings/dotnet/AIBridge/Models.cs @@ -0,0 +1,277 @@ +using System.Text.Json; +using System.Text.Json.Serialization; + +namespace AIBridge; + +// ============================================================================ +// 数据模型层 +// +// 对应 crates/aibridge-core/src/model/{chat,audio}.rs 的 serde struct。 +// 用 System.Text.Json 反序列化 FFI 返回的 JSON。字段名与 Rust serde 输出对齐 +// (Rust 默认 snake_case,C# 这里显式 JsonPropertyName 对齐)。 +// +// 只覆盖 hello world 涉及的能力:chat / chat_stream / speech。 +// 其余能力(image/video/embed/transcribe/list_models/list_voices)留待后续阶段。 +// ============================================================================ + +/// 对话请求(对应 core ChatRequest)。 +public sealed class ChatRequest +{ + /// 模型名称(如 "echo-chat"、"gpt-4o")。 + [JsonPropertyName("model")] + public string Model { get; set; } = string.Empty; + + /// 消息列表。 + [JsonPropertyName("messages")] + public List Messages { get; set; } = new(); + + /// 温度系数(可选)。 + [JsonPropertyName("temperature")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public double? Temperature { get; set; } + + /// 最大生成 token 数(可选)。 + [JsonPropertyName("max_tokens")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public uint? MaxTokens { get; set; } + + /// 是否流式(chat_stream 不依赖此字段,由调用方法决定)。 + [JsonPropertyName("stream")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingDefault)] + public bool Stream { get; set; } + + /// 构造请求。 + public ChatRequest(string model, IEnumerable messages) + { + Model = model; + Messages = messages.ToList(); + } + + /// 供 JsonSerializer 用,外部应使用带参构造。 + public ChatRequest() { } +} + +/// +/// 对话消息(对应 core ChatMessage 的 tagged enum)。 +/// +/// Rust 用 #[serde(tag="role", rename_all="lowercase")]: +/// 序列化为 {"role":"user","content":"..."} 等。C# 这里用扁平结构 + role 字段, +/// 仅覆盖 system/user 纯文本场景(hello world 足够),其余变体后续扩展。 +/// +public sealed class ChatMessage +{ + /// 角色:system / user / assistant / tool。 + [JsonPropertyName("role")] + public string Role { get; set; } = "user"; + + /// 消息内容(纯文本场景;多模态见 core UserContent::Parts,暂不支持)。 + [JsonPropertyName("content")] + public string Content { get; set; } = string.Empty; + + /// 创建 user 消息。 + public static ChatMessage User(string content) => new() + { + Role = "user", + Content = content, + }; + + /// 创建 system 消息。 + public static ChatMessage System(string content) => new() + { + Role = "system", + Content = content, + }; + + /// 创建 assistant 消息。 + public static ChatMessage Assistant(string content) => new() + { + Role = "assistant", + Content = content, + }; +} + +/// 对话完成结果(对应 core ChatCompletion)。 +public sealed class ChatCompletion +{ + [JsonPropertyName("id")] + public string Id { get; set; } = string.Empty; + + [JsonPropertyName("object")] + public string Object { get; set; } = "chat.completion"; + + [JsonPropertyName("created")] + public ulong Created { get; set; } + + [JsonPropertyName("model")] + public string Model { get; set; } = string.Empty; + + [JsonPropertyName("choices")] + public List Choices { get; set; } = new(); + + [JsonPropertyName("usage")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public ChatUsage? Usage { get; set; } +} + +/// 对话选项(对应 core ChatChoice)。 +public sealed class ChatChoice +{ + [JsonPropertyName("index")] + public uint Index { get; set; } + + [JsonPropertyName("message")] + public ChoiceMessage Message { get; set; } = new(); + + [JsonPropertyName("finish_reason")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? FinishReason { get; set; } +} + +/// 完成结果中的消息(对应 core ChoiceMessage)。 +public sealed class ChoiceMessage +{ + [JsonPropertyName("role")] + public string Role { get; set; } = "assistant"; + + [JsonPropertyName("content")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Content { get; set; } +} + +/// Token 使用统计(对应 core ChatUsage)。 +public sealed class ChatUsage +{ + [JsonPropertyName("prompt_tokens")] + public ulong PromptTokens { get; set; } + + [JsonPropertyName("completion_tokens")] + public ulong CompletionTokens { get; set; } + + [JsonPropertyName("total_tokens")] + public ulong TotalTokens { get; set; } +} + +/// 流式增量块(对应 core ChatCompletionChunk)。 +public sealed class ChatCompletionChunk +{ + [JsonPropertyName("id")] + public string Id { get; set; } = string.Empty; + + [JsonPropertyName("object")] + public string Object { get; set; } = "chat.completion.chunk"; + + [JsonPropertyName("created")] + public ulong Created { get; set; } + + [JsonPropertyName("model")] + public string Model { get; set; } = string.Empty; + + [JsonPropertyName("choices")] + public List Choices { get; set; } = new(); + + [JsonPropertyName("usage")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public ChatUsage? Usage { get; set; } +} + +/// 流式增量(对应 core ChatCompletionDelta)。 +public sealed class ChatCompletionDelta +{ + [JsonPropertyName("index")] + public uint Index { get; set; } + + [JsonPropertyName("delta")] + public DeltaMessage Delta { get; set; } = new(); + + [JsonPropertyName("finish_reason")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? FinishReason { get; set; } +} + +/// 流式增量消息(对应 core DeltaMessage)。 +public sealed class DeltaMessage +{ + [JsonPropertyName("role")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Role { get; set; } + + [JsonPropertyName("content")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Content { get; set; } +} + +/// 文字转语音请求(对应 core SpeechRequest)。 +public sealed class SpeechRequest +{ + [JsonPropertyName("model")] + public string Model { get; set; } = string.Empty; + + [JsonPropertyName("input")] + public string Input { get; set; } = string.Empty; + + /// 音色规格(core VoiceSpec,对象含 voices 数组)。 + [JsonPropertyName("voice")] + public VoiceSpec Voice { get; set; } = new(); + + /// + /// 音频输出格式("mp3" / "opus" / "aac" / "flac" / "wav" / "pcm")。 + /// 可空:为 null 时 Rust 端用默认 "mp3"(skip_serializing_if 语义)。 + /// + [JsonPropertyName("response_format")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? ResponseFormat { get; set; } + + /// 构造请求。 + public SpeechRequest(string model, string input, string voice) + { + Model = model; + Input = input; + Voice = VoiceSpec.Single(voice); + } + + /// 供 JsonSerializer 用。 + public SpeechRequest() { } +} + +/// 音色规格(对应 core VoiceSpec)。 +public sealed class VoiceSpec +{ + [JsonPropertyName("voices")] + public List Voices { get; set; } = new(); + + /// 单个音色。 + public static VoiceSpec Single(string voice) => new() { Voices = new List { voice } }; +} + +/// +/// 文字转语音结果元数据(对应 core SpeechResult,audio_data 不序列化)。 +/// 二进制音频通过 FFI 的 aibridge_bytes_t 单独返回,见 SpeechResult.AudioData。 +/// +public sealed class SpeechResult +{ + /// 音频二进制数据(来自 FFI aibridge_bytes_t,非 JSON 字段)。 + [JsonIgnore] + public byte[] AudioData { get; set; } = Array.Empty(); + + [JsonPropertyName("audio_url")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? AudioUrl { get; set; } + + [JsonPropertyName("audio_base64")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? AudioBase64 { get; set; } + + [JsonPropertyName("content_type")] + public string ContentType { get; set; } = "audio/mpeg"; + + [JsonPropertyName("format")] + public string Format { get; set; } = "mp3"; + + [JsonPropertyName("duration")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public double? Duration { get; set; } + + [JsonPropertyName("model")] + [JsonIgnore(Condition = JsonIgnoreCondition.WhenWritingNull)] + public string? Model { get; set; } +} diff --git a/bindings/dotnet/AIBridge/Native.cs b/bindings/dotnet/AIBridge/Native.cs new file mode 100644 index 0000000..9e0f87e --- /dev/null +++ b/bindings/dotnet/AIBridge/Native.cs @@ -0,0 +1,320 @@ +using System.Runtime.InteropServices; + +namespace AIBridge; + +// ============================================================================ +// P/Invoke 声明层 +// +// 对应 crates/aibridge-ffi/include/aibridge.h 的全部 extern "C" 函数。 +// 所有声明严格匹配头文件签名(opaque 指针 / char** / aibridge_bytes_t**)。 +// +// 安全要点(设计文档第 7 节 + FFI 遗留问题): +// 1. last_error 线程局部:调 FFI 失败后,必须在同一托管线程立即读取 +// aibridge_last_error() 转存为字符串,再抛异常。绝不可跨线程读取。 +// 2. stream_next 串行:同一 stream 句柄不可并发 next(由 ChatStreamAsync 保证)。 +// 3. aibridge_bytes_free / aibridge_string_free:必须调用,用 SafeHandle 兜底。 +// 4. client / stream 句柄必须 destroy,用 SafeHandle / IDisposable 兜底。 +// ============================================================================ + +/// +/// FFI 错误码常量(与 aibridge.h 的 #define 一一对应)。 +/// +internal static class AibridgeStatus +{ + public const int Ok = 0; + + // stream_next 专用返回值 + public const int StreamChunk = 0; // 拉到一个 chunk + public const int StreamEnd = 1; // 流正常结束 + + // 错误类别(负数) + public const int Authentication = -1; + public const int RateLimit = -2; + public const int Validation = -3; + public const int ModelNotFound = -4; + public const int Api = -5; + public const int Network = -6; + public const int Timeout = -7; + public const int UnsupportedCapability = -8; + public const int ProviderNotFound = -9; + public const int VoiceNotAvailable = -10; + public const int ServiceUnavailable = -11; + public const int Ffi = -100; // FFI 层通用错误(空指针、JSON 解析失败、panic) +} + +/// +/// 二进制缓冲结构(对应 aibridge.h 的 aibridge_bytes_t)。 +/// #[repr(C)] 保证 ptr + len 布局,C# 用 LayoutKind.Sequential 对齐。 +/// +[StructLayout(LayoutKind.Sequential)] +internal struct AibridgeBytes +{ + public IntPtr ptr; // const uint8_t*(Rust 分配) + public UIntPtr len; // size_t +} + +/// +/// Rust 分配的 C 字符串 SafeHandle。 +/// 包装 aibridge_string_free,保证即使异常也会释放,避免内存泄漏。 +/// +internal sealed class AibridgeStringHandle : SafeHandle +{ + public AibridgeStringHandle() : base(IntPtr.Zero, ownsHandle: true) { } + + public override bool IsInvalid => handle == IntPtr.Zero; + + protected override bool ReleaseHandle() + { + // 传 nullptr 是安全的 no-op,直接释放 + Native.aibridge_string_free(handle); + return true; + } + + /// 把句柄转为托管字符串(UTF-8 → string),并释放原生缓冲。 + public string? MarshalAndFree() + { + if (IsInvalid) return null; + // 先拷贝为托管字符串,再释放原生内存(避免悬垂指针) + string? s = Marshal.PtrToStringUTF8(handle); + Dispose(); + return s; + } +} + +/// +/// Rust 分配的二进制缓冲 SafeHandle。 +/// 包装 aibridge_bytes_free,保证二进制音频缓冲被释放。 +/// +internal sealed class AibridgeBytesHandle : SafeHandle +{ + public AibridgeBytesHandle() : base(IntPtr.Zero, ownsHandle: true) { } + + public override bool IsInvalid => handle == IntPtr.Zero; + + protected override bool ReleaseHandle() + { + Native.aibridge_bytes_free(handle); + return true; + } + + /// 把缓冲拷贝为 byte[] 并释放原生内存。 + public byte[] MarshalAndFree() + { + if (IsInvalid) return Array.Empty(); + // 先解引用结构拿到 ptr + len;用 try/finally 保证异常路径也释放原生内存 + // (new byte[] 在超大 len 时可能 OOM,此时仍需释放 aibridge_bytes_t) + byte[] data; + try + { + AibridgeBytes b = Marshal.PtrToStructure(handle); + // checked 防截断;len 超过 int.MaxValue 在 .NET byte[] 上限外,会抛异常但仍走 finally 释放 + int len = checked((int)b.len.ToUInt64()); + data = new byte[len]; + if (b.ptr != IntPtr.Zero && len > 0) + { + Marshal.Copy(b.ptr, data, 0, len); + } + } + finally + { + Dispose(); // 释放原生 aibridge_bytes_t(含内部 [u8]) + } + return data; + } +} + +/// +/// P/Invoke 入口声明。DllImport 使用 "aibridge"(不带 lib 前缀和扩展名), +/// 运行时由 NativeLibrary 按 OS 解析为 libaibridge.{dylib,so} / aibridge.dll。 +/// 动态库定位逻辑见 。 +/// +internal static class Native +{ + private const string LibName = "aibridge"; + + // 静态构造:注册动态库解析回调,支持 AIBRIDGE_LIB_PATH 环境变量与输出目录搜索。 + static Native() + { + NativeResolver.Register(LibName); + } + + // —— 生命周期 —— ------------------------------------------------------------ + + /// 创建客户端(对应 aibridge_client_new)。 + /// Provider 类型,UTF-8 C 字符串(如 "echo"、"openai")。 + /// ClientOptions 的 JSON,可为 null(默认配置)。 + /// 成功返 client 指针;失败返 IntPtr.Zero(读 last_error)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern IntPtr aibridge_client_new(byte[] provider, byte[]? configJson); + + /// 启动客户端(对应 aibridge_client_start)。 + /// 0 成功;负数为错误码。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern int aibridge_client_start(IntPtr client); + + /// 释放客户端句柄(对应 aibridge_client_destroy,nullptr 安全)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern void aibridge_client_destroy(IntPtr client); + + // —— 阻塞式调用 —— --------------------------------------------------------- + + /// 文本对话(对应 aibridge_client_chat)。 + /// ChatRequest 的 JSON。 + /// 输出 ChatCompletion 的 JSON,调用方 aibridge_string_free。 + /// 0 成功;负数为错误码。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern int aibridge_client_chat( + IntPtr client, + byte[] requestJson, + ref IntPtr outResponseJson); + + /// 文字转语音(对应 aibridge_client_speech,二进制走 aibridge_bytes_t)。 + /// SpeechRequest 的 JSON。 + /// 输出二进制音频缓冲(可能为 IntPtr.Zero)。 + /// 输出 SpeechResult(不含 audio_data)的 JSON。 + /// 0 成功;负数为错误码。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern int aibridge_client_speech( + IntPtr client, + byte[] requestJson, + ref IntPtr outAudio, + ref IntPtr outMetaJson); + + // —— 流式 —— ---------------------------------------------------------------- + + /// 创建流式 stream 句柄(对应 aibridge_client_chat_stream)。 + /// ChatRequest 的 JSON。 + /// 输出 stream 句柄。 + /// 0 成功;负数为错误码。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern int aibridge_client_chat_stream( + IntPtr client, + byte[] requestJson, + ref IntPtr outStream); + + /// 拉取下一个流式 chunk(阻塞,对应 aibridge_stream_next)。 + /// 0=chunk,1=EOF,负数=错误。outChunkJson 为 chunk JSON(需 aibridge_string_free)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern int aibridge_stream_next(IntPtr stream, ref IntPtr outChunkJson); + + /// 释放 stream 句柄(触发 Rust drop → tokio task abort,nullptr 安全)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern void aibridge_stream_destroy(IntPtr stream); + + // —— 错误与释放 —— ---------------------------------------------------------- + + /// 读取当前线程的 last_error(JSON 字符串,调用方不应释放)。 + /// 错误 JSON 指针;无错误返 IntPtr.Zero。线程局部,仅当前线程下一次 FFI 调用前有效。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern IntPtr aibridge_last_error(); + + /// 释放 Rust 分配的 C 字符串(nullptr 安全)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern void aibridge_string_free(IntPtr ptr); + + /// 释放 Rust 分配的二进制缓冲(nullptr 安全)。 + [DllImport(LibName, CallingConvention = CallingConvention.Cdecl)] + public static extern void aibridge_bytes_free(IntPtr ptr); +} + +/// +/// 动态库运行时解析器。 +/// +/// .NET 默认只在系统库目录和输出目录搜索 libaibridge。为方便开发期直接 +/// dotnet run(不打包),这里通过 NativeLibrary.SetDllImportResolver +/// 注入自定义解析:依次尝试 +/// 1) AIBRIDGE_LIB_PATH 环境变量指向的目录 +/// 2) 程序集输出目录(已通过 csproj CopyToOutputDirectory 拷贝过来) +/// 3) 仓库 target/debug 与 target/release +/// 找到后用 NativeLibrary.Load 载入并缓存句柄。 +/// +internal static class NativeResolver +{ + private static int _registered; // 0=未注册,1=已注册 + + public static void Register(string libraryName) + { + if (Interlocked.CompareExchange(ref _registered, 1, 0) != 0) return; + + // 解析回调签名:(string libName, Assembly asm, DllImportSearchPath? searchPath, IntPtr) => IntPtr + IntPtr Resolver(string lib, Assembly asm, DllImportSearchPath? search, IntPtr callers) + { + // 仅处理本绑定的库名(其它库走默认解析) + if (!string.Equals(lib, libraryName, StringComparison.OrdinalIgnoreCase)) + { + return IntPtr.Zero; + } + + string?[] candidates = + { + Environment.GetEnvironmentVariable("AIBRIDGE_LIB_PATH"), + AppContext.BaseDirectory, + }; + + // 候选目录优先级:环境变量 > 输出目录 + foreach (string? candidate in candidates) + { + if (string.IsNullOrEmpty(candidate)) continue; + + // 容错:candidate 可能是目录也可能是完整文件路径(用户误填) + string dir = candidate; + if (File.Exists(candidate)) + { + dir = Path.GetDirectoryName(candidate) ?? string.Empty; + } + + if (string.IsNullOrEmpty(dir)) continue; + string full = Path.Combine(dir, NativeFileName(libraryName)); + if (File.Exists(full)) + { + return NativeLibrary.Load(full, asm, DllImportSearchPath.Default); + } + } + + // 兜底:仓库 target/{debug,release}(相对输出目录回溯) + string repoRoot = FindRepoRoot(AppContext.BaseDirectory); + if (!string.IsNullOrEmpty(repoRoot)) + { + foreach (string profile in new[] { "debug", "release" }) + { + string full = Path.Combine(repoRoot, "target", profile, NativeFileName(libraryName)); + if (File.Exists(full)) + { + return NativeLibrary.Load(full, asm, DllImportSearchPath.Default); + } + } + } + + // 最后让 .NET 用默认搜索(系统库目录等) + return NativeLibrary.Load(libraryName, asm, DllImportSearchPath.Default); + } + + NativeLibrary.SetDllImportResolver(typeof(Native).Assembly, Resolver); + } + + /// 根据 OS 返回实际文件名(libaibridge.dylib / libaibridge.so / aibridge.dll)。 + private static string NativeFileName(string baseName) + { + if (RuntimeInformation.IsOSPlatform(OSPlatform.Windows)) + return baseName + ".dll"; + if (RuntimeInformation.IsOSPlatform(OSPlatform.OSX)) + return "lib" + baseName + ".dylib"; + return "lib" + baseName + ".so"; // Linux/Unix + } + + /// 从起始目录向上查找仓库根(含 target/ + Cargo.toml 的目录)。 + private static string FindRepoRoot(string start) + { + DirectoryInfo? dir = new(start); + while (dir != null) + { + if (Directory.Exists(Path.Combine(dir.FullName, "target")) + && File.Exists(Path.Combine(dir.FullName, "Cargo.toml"))) + { + return dir.FullName; + } + dir = dir.Parent; + } + return string.Empty; + } +} diff --git a/bindings/dotnet/build.sh b/bindings/dotnet/build.sh new file mode 100755 index 0000000..557d508 --- /dev/null +++ b/bindings/dotnet/build.sh @@ -0,0 +1,112 @@ +#!/usr/bin/env bash +# ============================================================================ +# AIBridge .NET 绑定构建与运行脚本 +# +# 用途: +# 1. cargo build -p aibridge-ffi 产出 libaibridge 动态库 +# 2. dotnet build 构建 C# 绑定(含 hello world) +# 3. dotnet run 跑通 hello world(echo 适配器,免认证) +# +# 前置:需安装 .NET 8 SDK。未装时本脚本会给出安装提示。 +# macOS: brew install --cask dotnet-sdk +# Linux: 见 https://learn.microsoft.com/dotnet/core/install/linux +# Windows: 见 https://dotnet.microsoft.com/download +# +# 用法: +# ./build.sh # 构建 ffi + dotnet 项目 +# ./build.sh run # 构建并运行 hello world +# ============================================================================ + +set -euo pipefail + +# 仓库根目录(脚本位于 bindings/dotnet/build.sh,回溯三级) +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/../.." && pwd)" +FFI_LIB_NAME="aibridge" # P/Invoke 库名(OS 加 lib 前缀/扩展名) +DOTNET_DIR="${REPO_ROOT}/bindings/dotnet/AIBridge" + +# 颜色输出 +info() { printf "\033[1;34m[INFO]\033[0m %s\n" "$*"; } +ok() { printf "\033[1;32m[OK]\033[0m %s\n" "$*"; } +warn() { printf "\033[1;33m[WARN]\033[0m %s\n" "$*"; } +fail() { printf "\033[1;31m[FAIL]\033[0m %s\n" "$*"; exit 1; } + +# ---------- Step 1: 构建 aibridge-ffi ---------- +build_ffi() { + info "构建 aibridge-ffi(cargo build -p aibridge-ffi)..." + (cd "${REPO_ROOT}" && cargo build -p aibridge-ffi) + + # 确认动态库产物 + local lib_file + if [[ "$(uname)" == "Darwin" ]]; then + lib_file="${REPO_ROOT}/target/debug/libaibridge.dylib" + elif [[ "$(uname)" == *MINGW* ]] || [[ "$(uname)" == *MSYS* ]]; then + lib_file="${REPO_ROOT}/target/debug/aibridge.dll" + else + lib_file="${REPO_ROOT}/target/debug/libaibridge.so" + fi + + [[ -f "${lib_file}" ]] || fail "未找到动态库产物: ${lib_file}" + ok "FFI 动态库就绪: ${lib_file}" +} + +# ---------- Step 2: 检查 dotnet ---------- +check_dotnet() { + if ! command -v dotnet &>/dev/null; then + warn "未检测到 dotnet SDK。请先安装 .NET 8 SDK:" + echo " macOS: brew install --cask dotnet-sdk" + echo " Linux: https://learn.microsoft.com/dotnet/core/install/linux" + echo " Windows: https://dotnet.microsoft.com/download" + echo "" + echo "代码已就绪(${DOTNET_DIR}/),装好 dotnet 后重跑本脚本即可。" + echo "当前状态:dotnet hello world 待 dotnet 环境验证。" + return 1 + fi + ok "dotnet 可用: $(dotnet --version)" +} + +# ---------- Step 3: dotnet build ---------- +build_dotnet() { + info "构建 .NET 项目(dotnet build)..." + (cd "${DOTNET_DIR}" && dotnet build) + ok ".NET 项目构建完成" +} + +# ---------- Step 4: dotnet run ---------- +run_hello() { + info "运行 hello world(dotnet run)..." + echo "----------------------------------------" + (cd "${DOTNET_DIR}" && dotnet run --no-build) + local rc=$? + echo "----------------------------------------" + if [[ ${rc} -eq 0 ]]; then + ok "hello world 运行成功(exit 0)" + else + fail "hello world 运行失败(exit ${rc})" + fi +} + +# ---------- 主流程 ---------- +main() { + info "AIBridge .NET 绑定构建脚本" + info "仓库根: ${REPO_ROOT}" + echo "" + + build_ffi + echo "" + + if ! check_dotnet; then + exit 1 + fi + echo "" + + build_dotnet + echo "" + + if [[ "${1:-}" == "run" ]]; then + run_hello + else + info "构建完成。运行 hello world 请执行: $0 run" + fi +} + +main "$@" From b81dab593c7595f53aa0cccb747b3e9141c9e592 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 12:15:03 +0800 Subject: [PATCH 12/55] =?UTF-8?q?chore:=20=E6=8F=90=E4=BA=A4=E9=98=B6?= =?UTF-8?q?=E6=AE=B50.6=20=E4=BA=94=E8=AF=AD=E8=A8=80=E7=BB=91=E5=AE=9A?= =?UTF-8?q?=E4=BE=9D=E8=B5=96=E9=94=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 5 +++++ 1 file changed, 5 insertions(+) diff --git a/Cargo.lock b/Cargo.lock index 5f4a70e..ebd4432 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -48,6 +48,7 @@ name = "aibridge-node" version = "2.0.0-alpha.1" dependencies = [ "aibridge-core", + "futures", "napi", "napi-build", "napi-derive", @@ -60,6 +61,8 @@ name = "aibridge-python" version = "2.0.0-alpha.1" dependencies = [ "aibridge-core", + "futures", + "once_cell", "pyo3", "serde_json", "tokio", @@ -809,6 +812,8 @@ dependencies = [ "napi-derive", "napi-sys", "once_cell", + "serde", + "serde_json", "tokio", ] From 60de645c0dffb0876b92b278b922d46c682850a4 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 12:27:21 +0800 Subject: [PATCH 13/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1.0=20OpenAI=20=E5=85=BC=E5=AE=B9=E9=80=82=E9=85=8D=E5=99=A8?= =?UTF-8?q?=E5=9C=B0=E5=9F=BA?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 实现 OpenAiCompatAdapter 作为 OpenAI 兼容协议(agnes/openai/azure 等约 80% 适配器)的共享地基, 子适配器(阶段1.1 openai / 1.2 agnes)可复用其 chat/chat_stream/image_generate/ embed/list_models 实现,仅 override base_url、provider_type、特有参数等差异。 主要内容: - OpenAiCompatAdapter struct:持有 HttpClient、ProviderConfig、provider_type、 provider_name、base_url、capabilities - 核心方法实现 OpenAI 兼容协议: - chat:POST /chat/completions,构造请求体(model/messages/temperature/ max_tokens/tools/stream 等),解析响应 → ChatCompletion - chat_stream:POST /chat/completions stream=true,自实现 SSE 行流解析器 (LinesStream)解析 data: → ChatStream,支持心跳/空行/[DONE] - image_generate:POST /images/generations → ImageResult - embed:POST /embeddings → EmbeddingResult - list_models:GET /models → Vec(含模型类型推断) - 参数映射:预置 openai_compatible_mapping()(用 ParameterMapping 机制,OpenAI 协议即通用名故透传) - 错误映射:map_api_error 提取 OpenAI error.message / retry_after, 按 HTTP 状态码分类(401→Authentication,429→RateLimit+retry_after, 404→ModelNotFound,400→Validation,5xx→Api) - 请求体构造:通用参数序列化 + extra 透传 + 流式附带 stream_options.include_usage - 设计为可复用:方法均为 pub,子适配器通过组合委托调用,Rust 无继承故采用 '地基结构 + trait 方法转发' 模式 单测(45 个,mockito mock HTTP): - chat 正常路径:解析 completion / 发送 temperature+max_tokens / extra 透传 / tool_calls 解析 - chat 错误路径:401/429(含 retry_after)/404/500/400/unsupported - chat_stream:SSE 多 chunk 解析 / 心跳空行 / stream=true + stream_options / 401 / 429 - image_generate:url 解析 / b64 解析 / 401 / 429 / unsupported - embed:多向量解析 / 单输入 / 401 / 404 / unsupported - list_models:正常 / 类型过滤 / 401 / 429 / 500 - 错误映射单元测试 + 参数映射测试 + base_url 兜底 + build_chat_body 测试 新增 dev-dependencies:mockito 1 adapters/mod.rs 导出 openai_compat 模块;暂不注册到工厂(阶段1.1/1.2 注册)。 验收:cargo test -p aibridge-core 235 通过(190 原有 + 45 新增), clippy --all-targets -D warnings 无警告,cargo build -p aibridge-ffi 通过, cargo fmt 已格式化。 --- Cargo.lock | 124 +- crates/aibridge-core/Cargo.toml | 5 + crates/aibridge-core/src/adapters/mod.rs | 3 + .../src/adapters/openai_compat.rs | 1919 +++++++++++++++++ 4 files changed, 2049 insertions(+), 2 deletions(-) create mode 100644 crates/aibridge-core/src/adapters/openai_compat.rs diff --git a/Cargo.lock b/Cargo.lock index ebd4432..dfc9e0b 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -20,6 +20,7 @@ dependencies = [ "base64", "bytes", "futures", + "mockito", "once_cell", "rand 0.8.6", "reqwest", @@ -118,6 +119,16 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "assert-json-diff" +version = "2.0.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "47e4f2b81832e72834d7518d8487a0396a28cc408186a2e8854c0f98011faf12" +dependencies = [ + "serde", + "serde_json", +] + [[package]] name = "async-stream" version = "0.3.6" @@ -266,6 +277,15 @@ version = "1.0.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1d07550c9036bf2ae0c684c4297d503f838287c83c53686d05370d0e139ae570" +[[package]] +name = "colored" +version = "3.1.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "faf9468729b8cbcea668e36183cb69d317348c2e08e994829fb56ebfdfbaac34" +dependencies = [ + "windows-sys 0.61.2", +] + [[package]] name = "convert_case" version = "0.6.0" @@ -449,6 +469,18 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "getrandom" +version = "0.3.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "899def5c37c4fd7b2664648c28120ecec138e4d395b459e5ca34f9cce2dd77fd" +dependencies = [ + "cfg-if", + "libc", + "r-efi 5.3.0", + "wasip2", +] + [[package]] name = "getrandom" version = "0.4.3" @@ -458,7 +490,7 @@ dependencies = [ "cfg-if", "js-sys", "libc", - "r-efi", + "r-efi 6.0.0", "rand_core 0.10.1", "wasm-bindgen", ] @@ -533,6 +565,12 @@ version = "1.10.1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "6dbf3de79e51f3d586ab4cb9d5c3e2c14aa28ed23d180cf89b4df0454a69cc87" +[[package]] +name = "httpdate" +version = "1.0.3" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "df3b46402a9d5adb4c86a0cf463f42e19994e3ee891101b1841f30a545cb49a9" + [[package]] name = "hyper" version = "1.10.1" @@ -547,6 +585,7 @@ dependencies = [ "http", "http-body", "httparse", + "httpdate", "itoa", "pin-project-lite", "smallvec", @@ -801,6 +840,31 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "mockito" +version = "1.7.2" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "90820618712cab19cfc46b274c6c22546a82affcb3c3bdf0f29e3db8e1bb92c0" +dependencies = [ + "assert-json-diff", + "bytes", + "colored", + "futures-core", + "http", + "http-body", + "http-body-util", + "hyper", + "hyper-util", + "log", + "pin-project-lite", + "rand 0.9.4", + "regex", + "serde_json", + "serde_urlencoded", + "similar", + "tokio", +] + [[package]] name = "napi" version = "2.16.17" @@ -1064,6 +1128,12 @@ dependencies = [ "proc-macro2", ] +[[package]] +name = "r-efi" +version = "5.3.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "69cdb34c158ceb288df11e18b4bd39de994f6657d83847bdffdbd7f346754b0f" + [[package]] name = "r-efi" version = "6.0.0" @@ -1077,10 +1147,20 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "5ca0ecfa931c29007047d1bc58e623ab12e5590e8c7cc53200d5202b69266d8a" dependencies = [ "libc", - "rand_chacha", + "rand_chacha 0.3.1", "rand_core 0.6.4", ] +[[package]] +name = "rand" +version = "0.9.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "44c5af06bb1b7d3216d91932aed5265164bf384dc89cd6ba05cf59a35f5f76ea" +dependencies = [ + "rand_chacha 0.9.0", + "rand_core 0.9.5", +] + [[package]] name = "rand" version = "0.10.2" @@ -1102,6 +1182,16 @@ dependencies = [ "rand_core 0.6.4", ] +[[package]] +name = "rand_chacha" +version = "0.9.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "d3022b5f1df60f26e1ffddd6c66e8aa15de382ae63b3a0c1bfc0e4d3e3f325cb" +dependencies = [ + "ppv-lite86", + "rand_core 0.9.5", +] + [[package]] name = "rand_core" version = "0.6.4" @@ -1111,6 +1201,15 @@ dependencies = [ "getrandom 0.2.17", ] +[[package]] +name = "rand_core" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "76afc826de14238e6e8c374ddcc1fa19e374fd8dd986b0d2af0d02377261d83c" +dependencies = [ + "getrandom 0.3.4", +] + [[package]] name = "rand_core" version = "0.10.1" @@ -1378,6 +1477,12 @@ dependencies = [ "libc", ] +[[package]] +name = "similar" +version = "2.7.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bbbb5d9659141646ae647b42fe094daf6c6192d1620870b449d9557f748b2daa" + [[package]] name = "slab" version = "0.4.12" @@ -1762,6 +1867,15 @@ version = "0.11.1+wasi-snapshot-preview1" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "ccf3ec651a847eb01de73ccad15eb7d99f80485de043efb2f370cd654f4ea44b" +[[package]] +name = "wasip2" +version = "1.0.4+wasi-0.2.12" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b67efb37e106e55ce722a510d6b5f9c17f083e5fc79afc2badeb12cc313d9487" +dependencies = [ + "wit-bindgen", +] + [[package]] name = "wasm-bindgen" version = "0.2.126" @@ -1959,6 +2073,12 @@ version = "1.0.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "0592e1c9d151f854e6fd382574c3a0855250e1d9b2f99d9281c6e6391af352f1" +[[package]] +name = "wit-bindgen" +version = "0.57.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1ebf944e87a7c253233ad6766e082e3cd714b5d03812acc24c318f549614536e" + [[package]] name = "writeable" version = "0.6.3" diff --git a/crates/aibridge-core/Cargo.toml b/crates/aibridge-core/Cargo.toml index 896a933..f8cb667 100644 --- a/crates/aibridge-core/Cargo.toml +++ b/crates/aibridge-core/Cargo.toml @@ -24,3 +24,8 @@ bytes.workspace = true base64 = "0.22" # 随机数(router.rs 的 round_robin/random/weighted 策略用) rand = "0.8" + +[dev-dependencies] +# HTTP mock,用于适配器单测(openai_compat 等) +mockito = "1" +tokio = { workspace = true, features = ["full", "test-util"] } diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index b45fd55..3c7d8cd 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -12,3 +12,6 @@ /// Echo(Mock)适配器:阶段 0.6 五语言管线验证用,不调网络返固定/回显响应 pub mod echo; + +/// OpenAI 兼容协议适配器地基:阶段 1.0 实现,为 openai/agnes 等子适配器提供共享基础 +pub mod openai_compat; diff --git a/crates/aibridge-core/src/adapters/openai_compat.rs b/crates/aibridge-core/src/adapters/openai_compat.rs new file mode 100644 index 0000000..4b53213 --- /dev/null +++ b/crates/aibridge-core/src/adapters/openai_compat.rs @@ -0,0 +1,1919 @@ +//! OpenAI 兼容协议适配器地基 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/openai.py` 与各 OpenAI 兼容适配器 +//! 的公共部分。14 个适配器中约 80% 是 OpenAI 兼容协议(agnes/openai/azure/ +//! 聚合平台/中文平台兼容部分),本模块抽出通用 HTTP 请求构造 + 响应解析 + +//! 参数映射 + 错误映射逻辑,子适配器(openai/agnes)只需 override base_url、 +//! provider_type、特有参数等差异,复用本模块的 chat/image/embed/list_models 实现。 +//! +//! 设计要点(与设计文档第 10 节一致): +//! - `OpenAiCompatAdapter` 持有 `HttpClient`、`ProviderConfig`、provider_type 等 +//! - 实现 `Adapter` trait 的核心方法:chat / chat_stream / image_generate / embed / list_models +//! - 预置 `OPENAI_COMPATIBLE_MAPPING` 通用参数映射(用 `ParameterMapping` 机制) +//! - 错误映射:OpenAI API 错误(HTTP status + error.message)→ AibridgeError +//! (401→Authentication,429→RateLimit,404→ModelNotFound,5xx→Api 等) +//! +//! 注意:本模块为地基,不注册到工厂;具体子适配器(openai/agnes)在阶段 1.1/1.2 实现。 + +use futures::stream::{StreamExt, TryStreamExt}; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; + +use crate::adapter::{Capabilities, CapabilitySet, ChatStream}; +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatRequest, + ChoiceMessage, DeltaMessage, +}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::image::{ImageData, ImageRequest, ImageResult}; +use crate::model::options::{ + EmbedRequest, EmbeddingItem, EmbeddingResult, EmbeddingUsage, EmbeddingVector, ParameterMapping, +}; +use crate::util; + +/// OpenAI 官方默认 Base URL +pub const DEFAULT_OPENAI_BASE_URL: &str = "https://api.openai.com/v1"; + +/// OpenAI 兼容协议的通用参数映射 +/// +/// 对应 Python v1 `OPENAI_COMPATIBLE_MAPPING`。 +/// OpenAI 协议本身即"通用参数名",因此映射表为空(保持原名透传)。 +/// 子适配器可在此基础上叠加自己的 rename_map。 +/// +/// 保留为常量是为了与设计文档 10.2 节"预置常量 OPENAI_COMPATIBLE_MAPPING"对齐, +/// 后续若发现部分兼容平台需要重命名(如 max_tokens → maxOutputTokens),可在此扩展。 +pub fn openai_compatible_mapping() -> ParameterMapping { + ParameterMapping::new() +} + +/// OpenAI 兼容适配器地基 +/// +/// 持有 HTTP 客户端与 Provider 配置,实现 OpenAI 兼容协议的核心能力。 +/// 子适配器(openai/agnes)通过组合(而非继承)复用本结构的方法: +/// 子适配器内部持有一个 `OpenAiCompatAdapter`,并把自身的 base_url、provider_type +/// 等差异传入;本结构的方法均为 `pub`,子适配器可直接调用。 +/// +/// Rust 无继承,故采用"地基结构 + trait 方法转发"模式: +/// 子适配器实现 `Adapter` trait 时,将 chat/image/embed/list_models 等委托给 +/// 内部的 `OpenAiCompatAdapter` 实例。 +pub struct OpenAiCompatAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// Provider 类型标识(如 "openai"、"agnes"),用于能力错误信息 + provider_type: String, + /// Provider 显示名称(如 "OpenAI") + provider_name: String, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl OpenAiCompatAdapter { + /// 创建 OpenAI 兼容适配器 + /// + /// - `config`:Provider 配置,base_url 为 None 时用 `default_base_url` 兜底 + /// - `provider_type`:Provider 类型标识 + /// - `provider_name`:Provider 显示名称 + /// - `default_base_url`:config.base_url 为空时的兜底 base_url + /// - `capabilities`:支持的能力集合 + pub fn new( + config: ProviderConfig, + provider_type: impl Into, + provider_name: impl Into, + default_base_url: &str, + capabilities: CapabilitySet, + ) -> Result { + let provider_type = provider_type.into(); + let provider_name = provider_name.into(); + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| default_base_url.to_string()); + + // 构造 HttpClient:把 base_url 透传,便于 post_json 等方法自动拼接相对路径 + let opts = crate::config::ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + provider_type, + provider_name, + base_url, + capabilities, + }) + } + + /// 用显式的 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http( + http: HttpClient, + config: ProviderConfig, + provider_type: impl Into, + provider_name: impl Into, + capabilities: CapabilitySet, + ) -> Self { + let base_url = config + .base_url + .clone() + .unwrap_or_else(|| DEFAULT_OPENAI_BASE_URL.to_string()); + Self { + http, + config, + provider_type: provider_type.into(), + provider_name: provider_name.into(), + base_url, + capabilities, + } + } + + /// Provider 类型标识 + pub fn provider_type_str(&self) -> &str { + &self.provider_type + } + + /// Provider 显示名称 + pub fn provider_name_str(&self) -> &str { + &self.provider_name + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// API key(可能为空,免费 provider 场景) + pub fn api_key(&self) -> Option<&str> { + self.config.api_key.as_deref() + } + + /// 支持的能力集合 + pub fn capabilities_set(&self) -> &CapabilitySet { + &self.capabilities + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), self.provider_type), + }) + } + } + + /// 发送带认证的 POST JSON 请求,并用 OpenAI 错误映射处理响应 + /// + /// 不依赖 HttpClient 的自动错误处理(其 from_http_status 不提取 OpenAI + /// error.message / retry_after),而是用 map_api_error 统一映射。 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带认证的 GET 请求,并用 OpenAI 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 OpenAI chat/completions 请求体 + /// + /// 将统一 `ChatRequest` 转为 OpenAI 协议的 JSON 请求体。 + /// 通用参数走 `OPENAI_COMPATIBLE_MAPPING`(当前为透传), + /// provider 特有参数走 `extra` 透传。 + fn build_chat_body(&self, req: &ChatRequest, stream: bool) -> Value { + // 序列化统一请求,得到基础字段 + let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); + // 强制覆盖 stream 标志(统一请求的 stream 字段默认 false,流式调用时需置 true) + if stream { + body["stream"] = json!(true); + // 流式场景附带 stream_options.include_usage,便于末尾拿到 usage 统计 + body["stream_options"] = json!({ "include_usage": true }); + } else if body + .get("stream") + .and_then(|v| v.as_bool()) + .unwrap_or(false) + { + // 非流式调用但请求体带了 stream=true,移除避免歧义 + body["stream"] = json!(false); + } + // extra 透传:合并到顶层(extra 字段本身不发送) + if let Some(obj) = body.as_object_mut() { + if let Some(extra) = obj.remove("extra") { + if let Some(extra_map) = extra.as_object() { + for (k, v) in extra_map { + obj.insert(k.clone(), v.clone()); + } + } + } + } + body + } + + /// 文本对话(非流式) + /// + /// POST /chat/completions,构造 OpenAI 请求体,解析响应 → ChatCompletion。 + pub async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = self.build_chat_body(&req, false); + let value = self.post_authed_json("chat/completions", &body).await?; + self.parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话 + /// + /// POST /chat/completions stream=true,解析 SSE 流 → ChatStream。 + /// SSE 格式:每行 `data: `,`data: [DONE]` 标记结束。 + pub async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = self.build_chat_body(&req, true); + let url = self.url("chat/completions"); + + // 流式请求需直接用 reqwest::Client(HttpClient 未暴露流式接口) + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(&body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + + let model = req.model.clone(); + // 按字节流读取,按行切分解析 SSE + // 把 reqwest::Error 统一转成 String,便于 LinesStream 跨错误类型复用 + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = LinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(以 ":" 开头的心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 去除 "data: " 前缀 + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + // 非 data 行,跳过 + continue; + }; + // 结束标记 + if data.trim() == "[DONE]" { + return; + } + // 解析 JSON + match serde_json::from_str::(data) { + Ok(v) => { + match Self::parse_chunk(&v, &model) { + Ok(Some(chunk)) => yield Ok(chunk), + Ok(None) => continue, + Err(e) => { + yield Err(e); + return; + } + } + } + Err(_) => { + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + continue; + } + } + } + // 流自然结束(未收到 [DONE])也视为正常 + }; + + Ok(stream.boxed()) + } + + /// 图像生成 + /// + /// POST /images/generations → ImageResult + pub async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let body = self.build_image_body(&req); + let value = self.post_authed_json("images/generations", &body).await?; + self.parse_image_result(&value, &req.model) + } + + /// 文本嵌入 + /// + /// POST /embeddings → EmbeddingResult + pub async fn embed(&self, req: EmbedRequest) -> Result { + self.ensure_capability(Capabilities::Embedding)?; + let body = self.build_embed_body(&req); + let value = self.post_authed_json("embeddings", &body).await?; + self.parse_embedding_result(&value, &req.model) + } + + /// 模型列表(实时拉取) + /// + /// GET /models → Vec,按 `filter` 过滤模型类型。 + pub async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("models").await?; + let models = self.parse_models(&value); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // ==================== 内部:请求体构造 ==================== + + /// 构造 OpenAI images/generations 请求体 + fn build_image_body(&self, req: &ImageRequest) -> Value { + let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); + // extra 透传 + if let Some(obj) = body.as_object_mut() { + if let Some(extra) = obj.remove("extra") { + if let Some(extra_map) = extra.as_object() { + for (k, v) in extra_map { + obj.insert(k.clone(), v.clone()); + } + } + } + } + body + } + + /// 构造 OpenAI embeddings 请求体 + fn build_embed_body(&self, req: &EmbedRequest) -> Value { + let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); + // extra 透传 + if let Some(obj) = body.as_object_mut() { + if let Some(extra) = obj.remove("extra") { + if let Some(extra_map) = extra.as_object() { + for (k, v) in extra_map { + obj.insert(k.clone(), v.clone()); + } + } + } + } + body + } + + // ==================== 内部:响应解析 ==================== + + /// 解析 OpenAI chat/completions 响应 → ChatCompletion + fn parse_chat_completion(&self, value: &Value, fallback_model: &str) -> Result { + let id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")); + let created = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + let model = value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()); + let object = value + .get("object") + .and_then(|v| v.as_str()) + .unwrap_or("chat.completion") + .to_string(); + let service_tier = value + .get("service_tier") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let system_fingerprint = value + .get("system_fingerprint") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let usage = value.get("usage").and_then(parse_usage); + + let choices = value + .get("choices") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .enumerate() + .map(|(i, c)| parse_choice(c, i)) + .collect() + }) + .unwrap_or_default(); + + Ok(ChatCompletion { + id, + object, + created, + model, + choices, + usage, + service_tier, + system_fingerprint, + }) + } + + /// 解析单个 SSE chunk(OpenAI 流式格式)→ Option + /// + /// 返回 None 表示该 chunk 无有效 choices(如纯 usage 块),调用方跳过。 + fn parse_chunk(value: &Value, fallback_model: &str) -> Result> { + let id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")); + let created = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + let model = value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()); + let object = value + .get("object") + .and_then(|v| v.as_str()) + .unwrap_or("chat.completion.chunk") + .to_string(); + let usage = value.get("usage").and_then(parse_usage); + + let choices: Vec = value + .get("choices") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .enumerate() + .map(|(i, c)| parse_delta(c, i)) + .collect() + }) + .unwrap_or_default(); + + // 无 choices 且无 usage 的空块跳过 + if choices.is_empty() && usage.is_none() { + return Ok(None); + } + Ok(Some(ChatCompletionChunk { + id, + object, + created, + model, + choices, + usage, + })) + } + + /// 解析 OpenAI images/generations 响应 → ImageResult + fn parse_image_result(&self, value: &Value, fallback_model: &str) -> Result { + let id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("img")); + let created = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + let model = value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()); + let object = value + .get("object") + .and_then(|v| v.as_str()) + .unwrap_or("image.generation") + .to_string(); + let data = value + .get("data") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_image_data).collect()) + .unwrap_or_default(); + Ok(ImageResult { + id, + object, + created, + model, + data, + }) + } + + /// 解析 OpenAI embeddings 响应 → EmbeddingResult + fn parse_embedding_result( + &self, + value: &Value, + fallback_model: &str, + ) -> Result { + let object = value + .get("object") + .and_then(|v| v.as_str()) + .unwrap_or("list") + .to_string(); + let model = value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()); + let data = value + .get("data") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_embedding_item).collect()) + .unwrap_or_default(); + let usage = value.get("usage").and_then(parse_embed_usage); + Ok(EmbeddingResult { + object, + data, + model, + usage, + }) + } + + /// 解析 OpenAI /models 响应 → Vec + /// + /// OpenAI 返回 `{"data": [{"id": "...", "created": ..., "owned_by": "..."}]}` + fn parse_models(&self, value: &Value) -> Vec { + let arr = value.get("data").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .map(|m| { + let id = m + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let model_type = infer_model_type(&id); + ModelInfo { + name: id.clone(), + id, + model_type, + provider: self.provider_type.clone(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: None, + created: m.get("created").and_then(|v| v.as_u64()), + } + }) + .collect(), + None => Vec::new(), + } + } + + // ==================== 内部:错误映射 ==================== + + /// 将 OpenAI API 错误响应映射为 AibridgeError + /// + /// 优先提取 OpenAI 错误体 `{"error": {"message": "..."}}` 中的 message, + /// 再按 HTTP 状态码分类: + /// - 401/403 → Authentication + /// - 429 → RateLimit(尝试从 Retry-After / 错误体提取 retry_after) + /// - 404 → ModelNotFound + /// - 400 → Validation + /// - 4xx(其他)→ Api + /// - 5xx → Api + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + // 尝试解析 OpenAI 错误结构 {"error": {"message": "...", "type": "..."}} + let message = parse_error_message(body, status); + match status { + 401 | 403 => AibridgeError::Authentication { message }, + 429 => { + let retry_after = parse_retry_after(body); + AibridgeError::RateLimit { + message, + retry_after, + } + } + 400 => AibridgeError::Validation { + message, + details: serde_json::json!({ "status_code": status, "response": body }), + }, + 404 => AibridgeError::ModelNotFound { model: message }, + s if (400..500).contains(&s) => AibridgeError::Api { status: s, message }, + s if (500..600).contains(&s) => AibridgeError::Api { status: s, message }, + s => AibridgeError::Api { status: s, message }, + } + } +} + +/// 解析 OpenAI 错误体中的 message 字段 +/// +/// OpenAI 错误格式:`{"error": {"message": "...", "type": "...", "code": "..."}}` +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 部分兼容平台直接用顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +/// 从错误体尝试解析 retry_after(秒) +/// +/// OpenAI 限流响应偶尔携带 `error.retry_after` 或顶层 `retry_after`。 +fn parse_retry_after(body: &str) -> Option { + let v = serde_json::from_str::(body).ok()?; + v.get("error") + .and_then(|e| e.get("retry_after")) + .and_then(|r| r.as_f64()) + .or_else(|| v.get("retry_after").and_then(|r| r.as_f64())) +} + +/// 解析 usage 统计 +fn parse_usage(v: &Value) -> Option { + let prompt = v.get("prompt_tokens").and_then(|x| x.as_u64())?; + let completion = v + .get("completion_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(0); + let total = v + .get("total_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(prompt + completion); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: total, + }) +} + +/// 解析单个 choice → ChatChoice +fn parse_choice(c: &Value, index: usize) -> ChatChoice { + let idx = c + .get("index") + .and_then(|v| v.as_u64()) + .map(|i| i as u32) + .unwrap_or(index as u32); + let finish_reason = c + .get("finish_reason") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let message = c.get("message").cloned().unwrap_or(Value::Null); + let role = message + .get("role") + .and_then(|v| v.as_str()) + .unwrap_or("assistant") + .to_string(); + let content = message + .get("content") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let tool_calls = message + .get("tool_calls") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_tool_call).collect()); + ChatChoice { + index: idx, + message: ChoiceMessage { + role, + content, + tool_calls, + }, + finish_reason, + } +} + +/// 解析单个流式 delta → ChatCompletionDelta +fn parse_delta(c: &Value, index: usize) -> ChatCompletionDelta { + let idx = c + .get("index") + .and_then(|v| v.as_u64()) + .map(|i| i as u32) + .unwrap_or(index as u32); + let finish_reason = c + .get("finish_reason") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let delta = c.get("delta").cloned().unwrap_or(Value::Null); + let role = delta + .get("role") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let content = delta + .get("content") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let tool_calls = delta + .get("tool_calls") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_tool_call).collect()); + ChatCompletionDelta { + index: idx, + delta: DeltaMessage { + role, + content, + tool_calls, + }, + finish_reason, + } +} + +/// 解析单个 ToolCall(OpenAI 格式) +fn parse_tool_call(v: &Value) -> crate::model::options::ToolCall { + let id = v + .get("id") + .and_then(|x| x.as_str()) + .unwrap_or("") + .to_string(); + let tool_type = v + .get("type") + .and_then(|x| x.as_str()) + .unwrap_or("function") + .to_string(); + let function = v.get("function").cloned().unwrap_or(Value::Null); + let name = function + .get("name") + .and_then(|x| x.as_str()) + .unwrap_or("") + .to_string(); + let arguments = function + .get("arguments") + .and_then(|x| x.as_str()) + .unwrap_or("{}") + .to_string(); + crate::model::options::ToolCall { + id, + tool_type, + function: crate::model::options::ToolCallFunction { name, arguments }, + } +} + +/// 解析单个 ImageData +fn parse_image_data(v: &Value) -> ImageData { + ImageData { + url: v.get("url").and_then(|x| x.as_str()).map(str::to_owned), + b64_json: v + .get("b64_json") + .and_then(|x| x.as_str()) + .map(str::to_owned), + revised_prompt: v + .get("revised_prompt") + .and_then(|x| x.as_str()) + .map(str::to_owned), + } +} + +/// 解析单个 EmbeddingItem +fn parse_embedding_item(v: &Value) -> EmbeddingItem { + let index = v + .get("index") + .and_then(|x| x.as_u64()) + .map(|i| i as u32) + .unwrap_or(0); + let embedding = if let Some(arr) = v.get("embedding").and_then(|x| x.as_array()) { + EmbeddingVector::Float(arr.iter().filter_map(|x| x.as_f64()).collect()) + } else if let Some(s) = v.get("embedding").and_then(|x| x.as_str()) { + EmbeddingVector::Base64(s.to_string()) + } else { + EmbeddingVector::Float(Vec::new()) + }; + EmbeddingItem { + object: "embedding".to_string(), + index, + embedding, + } +} + +/// 解析嵌入 usage +fn parse_embed_usage(v: &Value) -> Option { + let prompt = v.get("prompt_tokens").and_then(|x| x.as_u64())?; + let total = v + .get("total_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(prompt); + Some(EmbeddingUsage { + prompt_tokens: prompt, + total_tokens: total, + }) +} + +// ==================== SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器 +/// +/// reqwest 的 `bytes_stream` 返回字节 chunk,需自行按 `\n` 切分。 +/// 本结构维护一个未完成行的缓冲区,逐 chunk 拼接出完整行。 +/// 泛型 `S` 须为 `Stream, String>>`(错误统一转 String 便于复用)。 +struct LinesStream { + inner: S, + buffer: Vec, +} + +impl LinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for LinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + // 先看缓冲区是否已有完整行 + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + // 去掉末尾 \n 与可能的 \r + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + // 缓冲区无完整行,拉取下一 chunk + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + // 继续循环,尝试从缓冲区切出行 + } + Poll::Ready(None) => { + // 流结束,把缓冲区剩余内容作为最后一行返回 + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +// ==================== 测试辅助序列化结构 ==================== + +/// OpenAI /models 响应的最小解析结构(测试用,验证反序列化对齐协议) +#[allow(dead_code)] +#[derive(Debug, Deserialize, Serialize)] +struct OpenAiModelsResponse { + object: String, + data: Vec, +} + +#[allow(dead_code)] +#[derive(Debug, Deserialize, Serialize)] +struct OpenAiModelEntry { + id: String, + object: String, + created: u64, + owned_by: String, +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::image::ImageRequest; + use crate::model::options::EmbedInput; + use mockito::Server; + use std::collections::HashMap; + + /// 构造测试用 OpenAiCompatAdapter(指向 mockito server) + fn make_adapter(server: &Server, caps: CapabilitySet) -> OpenAiCompatAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("openai", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + OpenAiCompatAdapter::with_http(http, config, "openai", "OpenAI", caps) + } + + /// 全能力集合(chat + chat_stream + image_generate + embedding) + fn full_caps() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::ToolCall); + caps + } + + // ============ chat 正常路径 ============ + + #[tokio::test] + async fn chat_success_parses_completion() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }); + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.model, "gpt-4o"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_sends_temperature_and_max_tokens() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gpt-4o", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_passes_extra_params_through() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gpt-4o", + "custom_param": "custom_value" + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_with_tool_calls_parses() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "chatcmpl-2", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": { + "role": "assistant", + "content": null, + "tool_calls": [{ + "id": "call_1", + "type": "function", + "function": {"name": "get_weather", "arguments": "{\"city\":\"Beijing\"}"} + }] + }, + "finish_reason": "tool_calls" + }] + }); + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("weather?")]).build(); + let resp = adapter.chat(req).await.unwrap(); + let tool_calls = resp.choices[0] + .message + .tool_calls + .as_ref() + .expect("应有 tool_calls"); + assert_eq!(tool_calls.len(), 1); + assert_eq!(tool_calls[0].id, "call_1"); + assert_eq!(tool_calls[0].function.name, "get_weather"); + assert_eq!(tool_calls[0].function.arguments, "{\"city\":\"Beijing\"}"); + } + + // ============ chat 错误路径 ============ + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body( + json!({"error": {"message": "Rate limit exceeded", "retry_after": 1.5}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::RateLimit { retry_after, .. } => { + assert_eq!(retry_after, Some(1.5)); + } + _ => panic!("应为 RateLimit"), + } + } + + #[tokio::test] + async fn chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body( + json!({"error": {"message": "The model 'gpt-x' does not exist"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-x", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(500) + .with_body(json!({"error": {"message": "Internal server error"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn chat_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(400) + .with_body(json!({"error": {"message": "max_tokens is invalid"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("max_tokens")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn chat_unsupported_capability_returns_error() { + // 不支持 Chat 能力 + let server = Server::new_async().await; + let adapter = make_adapter(&server, CapabilitySet::new()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ chat_stream 正常 + 错误路径 ============ + + #[tokio::test] + async fn chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + // 构造 SSE 响应:3 个 data 行 + [DONE] + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Missing) + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + // 第 1 块 role + assert_eq!( + chunks[0].choices[0].delta.role.as_deref(), + Some("assistant") + ); + // 拼接内容 + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + // 第 3 块 finish_reason + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_stream_handles_heartbeat_and_empty_lines() { + let mut server = Server::new_async().await; + let sse = ": heartbeat\n\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hi\"},\"finish_reason\":null}]}\n\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 1); + assert_eq!(chunks[0].choices[0].delta.content.as_deref(), Some("hi")); + } + + #[tokio::test] + async fn chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn chat_stream_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::RateLimit { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn chat_stream_sends_stream_true() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "stream": true, + "stream_options": {"include_usage": true} + }))) + .with_status(200) + .with_body("data: [DONE]\n") + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + // ============ image_generate 正常 + 错误路径 ============ + + #[tokio::test] + async fn image_generate_success_parses_url() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{ + "url": "https://example.com/img.png", + "revised_prompt": "a cute cat" + }] + }); + server + .mock("POST", "/images/generations") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ImageRequest::builder("dall-e-3", "a cat") + .size("1024x1024") + .build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://example.com/img.png") + ); + assert_eq!(resp.data[0].revised_prompt.as_deref(), Some("a cute cat")); + } + + #[tokio::test] + async fn image_generate_success_parses_b64() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{"b64_json": "aGVsbG8="}] + }); + server + .mock("POST", "/images/generations") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + } + + #[tokio::test] + async fn image_generate_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn image_generate_error_429() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn image_generate_unsupported_capability() { + let server = Server::new_async().await; + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + let adapter = make_adapter(&server, caps); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ embed 正常 + 错误路径 ============ + + #[tokio::test] + async fn embed_success_parses_vectors() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}, + {"object": "embedding", "index": 1, "embedding": [0.4, 0.5, 0.6]} + ], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 4, "total_tokens": 4} + }); + server + .mock("POST", "/embeddings") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Multiple(vec!["a".into(), "b".into()]), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 2); + assert_eq!(resp.data[0].index, 0); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2, 0.3]); + } else { + panic!("应为 Float 向量"); + } + assert_eq!(resp.usage.as_ref().unwrap().prompt_tokens, 4); + } + + #[tokio::test] + async fn embed_single_input() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/embeddings") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "text-embedding-3-small", + "input": "hello" + }))) + .with_status(200) + .with_body( + json!({ + "object": "list", + "data": [{"object":"embedding","index":0,"embedding":[0.1]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn embed_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(404) + .with_body(json!({"error": {"message": "model not found"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let req = EmbedRequest { + model: "embed-x".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn embed_unsupported_capability() { + let server = Server::new_async().await; + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + let adapter = make_adapter(&server, caps); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ list_models 正常 + 错误路径 ============ + + #[tokio::test] + async fn list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"id": "gpt-4o", "object": "model", "created": 1700000000, "owned_by": "openai"}, + {"id": "dall-e-3", "object": "model", "created": 1700000000, "owned_by": "openai"}, + {"id": "whisper-1", "object": "model", "created": 1700000000, "owned_by": "openai"} + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + // 类型推断 + assert_eq!(models[0].id, "gpt-4o"); + assert_eq!(models[0].model_type, ModelType::Chat); + assert_eq!(models[1].model_type, ModelType::Image); + assert_eq!(models[2].model_type, ModelType::Audio); + // provider 字段填充 + assert_eq!(models[0].provider, "openai"); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "gpt-4o", "object": "model", "created": 1, "owned_by": "openai"}, + {"id": "dall-e-3", "object": "model", "created": 1, "owned_by": "openai"} + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, full_caps()); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn list_models_error_500() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server, full_caps()); + let err = adapter.list_models(None).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn map_api_error_401_with_message() { + let body = json!({"error": {"message": "Invalid API key"}}).to_string(); + let err = OpenAiCompatAdapter::map_api_error(401, &body); + match err { + AibridgeError::Authentication { message } => { + assert_eq!(message, "Invalid API key"); + } + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn map_api_error_429_extracts_retry_after() { + let body = json!({"error": {"message": "slow down", "retry_after": 2.0}}).to_string(); + let err = OpenAiCompatAdapter::map_api_error(429, &body); + match err { + AibridgeError::RateLimit { retry_after, .. } => { + assert_eq!(retry_after, Some(2.0)); + } + _ => panic!("应为 RateLimit"), + } + } + + #[test] + fn map_api_error_404_uses_message_as_model() { + let body = json!({"error": {"message": "model gpt-x not found"}}).to_string(); + let err = OpenAiCompatAdapter::map_api_error(404, &body); + match err { + AibridgeError::ModelNotFound { model } => { + assert!(model.contains("gpt-x")); + } + _ => panic!("应为 ModelNotFound"), + } + } + + #[test] + fn map_api_error_500_is_api() { + let err = OpenAiCompatAdapter::map_api_error(503, "service unavailable"); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 503), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_body_falls_back_to_http_status() { + let err = OpenAiCompatAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => { + assert!(message.contains("502")); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_400_is_validation() { + let body = json!({"error": {"message": "bad param"}}).to_string(); + let err = OpenAiCompatAdapter::map_api_error(400, &body); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + // ============ 参数映射 ============ + + #[test] + fn openai_compatible_mapping_is_passthrough() { + let pm = openai_compatible_mapping(); + let mut params = HashMap::new(); + params.insert("max_tokens".to_string(), json!(1000)); + params.insert("temperature".to_string(), json!(0.7)); + let result = pm.apply(¶ms); + // OpenAI 兼容映射为空表,参数原名透传 + assert_eq!( + result.get("max_tokens").and_then(|v| v.as_i64()), + Some(1000) + ); + assert_eq!( + result.get("temperature").and_then(|v| v.as_f64()), + Some(0.7) + ); + } + + // ============ 辅助方法 ============ + + #[tokio::test] + async fn base_url_uses_config_when_provided() { + // new() 应使用 config.base_url,而非默认值 + let config = ProviderConfig::from_options( + "openai", + ClientOptions::builder() + .api_key("k") + .base_url("https://custom.example.com/v1") + .build(), + ); + let adapter = OpenAiCompatAdapter::new( + config, + "openai", + "OpenAI", + DEFAULT_OPENAI_BASE_URL, + full_caps(), + ) + .unwrap(); + assert_eq!(adapter.base_url(), "https://custom.example.com/v1"); + } + + #[tokio::test] + async fn base_url_falls_back_to_default_when_missing() { + let config = + ProviderConfig::from_options("openai", ClientOptions::builder().api_key("k").build()); + let adapter = OpenAiCompatAdapter::new( + config, + "openai", + "OpenAI", + DEFAULT_OPENAI_BASE_URL, + full_caps(), + ) + .unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_OPENAI_BASE_URL); + } + + #[test] + fn ensure_capability_blocks_unsupported() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("openai", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = + OpenAiCompatAdapter::with_http(http, config, "openai", "OpenAI", CapabilitySet::new()); + let err = adapter.ensure_capability(Capabilities::Chat).unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[test] + fn ensure_capability_allows_supported() { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("openai", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = OpenAiCompatAdapter::with_http(http, config, "openai", "OpenAI", caps); + assert!(adapter.ensure_capability(Capabilities::Chat).is_ok()); + } + + #[test] + fn build_chat_body_includes_messages_and_model() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("openai", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = OpenAiCompatAdapter::with_http(http, config, "openai", "OpenAI", full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .build(); + let body = adapter.build_chat_body(&req, false); + assert_eq!(body["model"], "gpt-4o"); + assert_eq!(body["temperature"], 0.7); + assert!(body.get("messages").is_some()); + // 非 stream 模式不应注入 stream_options + assert!(body.get("stream_options").is_none()); + } + + #[test] + fn build_chat_body_stream_adds_stream_options() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("openai", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = OpenAiCompatAdapter::with_http(http, config, "openai", "OpenAI", full_caps()); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let body = adapter.build_chat_body(&req, true); + assert_eq!(body["stream"], true); + assert_eq!(body["stream_options"]["include_usage"], true); + } + + #[test] + fn parse_chunk_skips_empty_choices_without_usage() { + let value = + json!({"id":"x","object":"chat.completion.chunk","created":1,"model":"m","choices":[]}); + let result = OpenAiCompatAdapter::parse_chunk(&value, "m").unwrap(); + assert!(result.is_none()); + } + + #[test] + fn parse_chunk_with_usage_returns_some() { + let value = json!({ + "id":"x","object":"chat.completion.chunk","created":1,"model":"m", + "choices":[], + "usage":{"prompt_tokens":5,"completion_tokens":2,"total_tokens":7} + }); + let result = OpenAiCompatAdapter::parse_chunk(&value, "m").unwrap(); + assert!(result.is_some()); + let chunk = result.unwrap(); + assert_eq!(chunk.usage.as_ref().unwrap().total_tokens, 7); + } +} From d9a66e31adf8728c940fa0e1f10adb1c0b6278da Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 19:41:19 +0800 Subject: [PATCH 14/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1.1=20openai=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/openai.rs | 665 ++++++++++++++++++++ 1 file changed, 665 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/openai.rs diff --git a/crates/aibridge-core/src/adapters/openai.rs b/crates/aibridge-core/src/adapters/openai.rs new file mode 100644 index 0000000..5eae2ad --- /dev/null +++ b/crates/aibridge-core/src/adapters/openai.rs @@ -0,0 +1,665 @@ +//! OpenAI 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/openai.py`。 +//! +//! OpenAI 是 OpenAI 兼容协议的本源,本适配器不重写任何 HTTP/解析逻辑, +//! 而是组合(委托)[`crate::adapters::openai_compat::OpenAiCompatAdapter`] 地基: +//! Rust 无继承,用"持有地基 + trait 方法转发"模式复用 80% 通用代码。 +//! +//! 与 Python 老版的行为对照: +//! - `provider_type = "openai"`、`provider_name = "OpenAI"` +//! - 默认 `base_url = https://api.openai.com/v1`(可被 config.base_url 覆盖) +//! - 能力:Chat / ChatStream / ImageGenerate / Embedding +//! - `list_models` 实时拉取 `GET /models` +//! - 不支持的能力(video / transcribe / speech / list_voices)走 trait 默认实现, +//! 返 `UnsupportedCapability`(与 Python `video_create` 抛错一致) +//! - `requires_api_key = true` + +use async_trait::async_trait; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::{OpenAiCompatAdapter, DEFAULT_OPENAI_BASE_URL}; +use crate::config::ProviderConfig; +use crate::error::Result; +use crate::model::chat::{ChatCompletion, ChatRequest}; +use crate::model::common::{ModelInfo, ModelType}; +use crate::model::image::{ImageRequest, ImageResult}; +use crate::model::options::{EmbedRequest, EmbeddingResult}; + +/// OpenAI 适配器 +/// +/// 组合持有 [`OpenAiCompatAdapter`] 地基,把 `Adapter` trait 的核心方法 +/// (chat / chat_stream / image_generate / embed / list_models)委托给地基实现。 +/// 本结构仅负责:provider 元信息、能力集合声明、构造地基。 +/// +/// 构造时即建立 HTTP 客户端(地基 `new` 内部完成),`start` / `close` 为空操作。 +pub struct OpenAiAdapter { + /// OpenAI 兼容协议地基,复用 chat/image/embed/list_models 实现 + compat: OpenAiCompatAdapter, +} + +impl OpenAiAdapter { + /// 创建 OpenAI 适配器 + /// + /// - `config.base_url` 为空时回退到 `https://api.openai.com/v1` + /// - 内部构造 `OpenAiCompatAdapter`(含 HTTP 客户端、连接池、错误映射) + pub fn new(config: ProviderConfig) -> Result { + let caps = Self::capabilities_set(); + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_OPENAI_BASE_URL, + caps, + )?; + Ok(Self { compat }) + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "openai"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "OpenAI"; + + /// 支持的能力集合 + /// + /// 对应 Python v1 `OpenAIAdapter.supported_capabilities` 的核心子集 + /// (chat / chat_stream / image_generate / embedding)。 + /// video / audio 不在 OpenAI 本适配器能力内,走 trait 默认实现返 UnsupportedCapability。 + fn capabilities_set() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::Embedding); + caps + } +} + +#[async_trait] +impl Adapter for OpenAiAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + // 返回新集合,避免外部修改内部状态(不可变原则) + Self::capabilities_set() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new() 时已构造,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 走 Drop 释放,无需显式关闭 + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 图像生成:委托地基 `POST /images/generations` + async fn image_generate(&self, req: ImageRequest) -> Result { + self.compat.image_generate(req).await + } + + /// 文本嵌入:委托地基 `POST /embeddings` + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // 其余方法(video_create / video_poll / transcribe / speech / list_voices) + // 走 trait 默认实现,返 UnsupportedCapability,与 Python 老版行为一致。 +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::error::AibridgeError; + use crate::model::chat::ChatMessage; + use crate::model::options::{EmbedInput, EmbeddingVector}; + use futures::stream::StreamExt; + use mockito::Server; + use serde_json::json; + use std::collections::HashMap; + + /// 构造指向 mockito server 的 OpenAiAdapter(base_url 注入 mock 地址) + fn make_adapter(server: &Server) -> OpenAiAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("openai", opts); + OpenAiAdapter::new(config).expect("OpenAiAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 OpenAiAdapter(用于不发请求的元信息/能力测试) + fn make_adapter_no_server() -> OpenAiAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url("https://api.openai.com/v1") + .build(); + let config = ProviderConfig::from_options("openai", opts); + OpenAiAdapter::new(config).expect("OpenAiAdapter 构造应成功") + } + + // ============ 元信息 ============ + + #[test] + fn provider_type_and_name_match_python() { + let adapter = make_adapter_no_server(); + assert_eq!(adapter.provider_type(), "openai"); + assert_eq!(adapter.provider_name(), "OpenAI"); + } + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_core_set() { + let adapter = make_adapter_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::Embedding)); + // video / audio 不应声明 + assert!(!caps.contains(&Capabilities::VideoGenerate)); + assert!(!caps.contains(&Capabilities::AudioSpeech)); + } + + #[test] + fn base_url_defaults_to_openai_when_missing() { + // 不提供 base_url,应回退到 OpenAI 官方默认值 + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("openai", opts); + let adapter = OpenAiAdapter::new(config).unwrap(); + // 地基的 base_url 暴露为 pub,可直接读取 + assert_eq!(adapter.compat.base_url(), DEFAULT_OPENAI_BASE_URL); + } + + #[test] + fn base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.openai-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("openai", opts); + let adapter = OpenAiAdapter::new(config).unwrap(); + assert_eq!( + adapter.compat.base_url(), + "https://custom.openai-proxy.com/v1" + ); + } + + // ============ chat 正常路径 ============ + + #[tokio::test] + async fn chat_success_returns_completion() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }); + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.model, "gpt-4o"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_sends_bearer_auth_and_model() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gpt-4o", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + // ============ chat 错误路径 ============ + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body( + json!({"error": {"message": "The model 'gpt-x' does not exist"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-x", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + // ============ chat_stream ============ + + #[tokio::test] + async fn chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + // ============ image_generate ============ + + #[tokio::test] + async fn image_generate_success_parses_url() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{ + "url": "https://example.com/img.png", + "revised_prompt": "a cute cat" + }] + }); + server + .mock("POST", "/images/generations") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat") + .size("1024x1024") + .build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://example.com/img.png") + ); + assert_eq!(resp.data[0].revised_prompt.as_deref(), Some("a cute cat")); + } + + #[tokio::test] + async fn image_generate_success_parses_b64() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{"b64_json": "aGVsbG8="}] + }); + server + .mock("POST", "/images/generations") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + } + + #[tokio::test] + async fn image_generate_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ embed ============ + + #[tokio::test] + async fn embed_success_parses_vectors() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}, + {"object": "embedding", "index": 1, "embedding": [0.4, 0.5, 0.6]} + ], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 4, "total_tokens": 4} + }); + server + .mock("POST", "/embeddings") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Multiple(vec!["a".into(), "b".into()]), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 2); + assert_eq!(resp.data[0].index, 0); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2, 0.3]); + } else { + panic!("应为 Float 向量"); + } + assert_eq!(resp.usage.as_ref().unwrap().prompt_tokens, 4); + } + + #[tokio::test] + async fn embed_single_input_sends_string() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/embeddings") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "text-embedding-3-small", + "input": "hello" + }))) + .with_status(200) + .with_body( + json!({ + "object": "list", + "data": [{"object":"embedding","index":0,"embedding":[0.1]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"id": "gpt-4o", "object": "model", "created": 1700000000, "owned_by": "openai"}, + {"id": "dall-e-3", "object": "model", "created": 1700000000, "owned_by": "openai"}, + {"id": "whisper-1", "object": "model", "created": 1700000000, "owned_by": "openai"} + ] + }); + server + .mock("GET", "/models") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + assert_eq!(models[0].id, "gpt-4o"); + assert_eq!(models[0].model_type, ModelType::Chat); + assert_eq!(models[1].model_type, ModelType::Image); + assert_eq!(models[2].model_type, ModelType::Audio); + // provider 字段填充为 "openai" + assert_eq!(models[0].provider, "openai"); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "gpt-4o", "object": "model", "created": 1, "owned_by": "openai"}, + {"id": "dall-e-3", "object": "model", "created": 1, "owned_by": "openai"} + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn list_models_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ 不支持的能力走默认实现 ============ + + #[tokio::test] + async fn video_create_returns_unsupported() { + // 与 Python 老版 `video_create` 抛 UnsupportedCapabilityError 行为一致 + let adapter = make_adapter_no_server(); + let req = crate::model::video::VideoRequest::builder("gpt-4o", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn speech_returns_unsupported() { + let adapter = make_adapter_no_server(); + let req = crate::model::audio::SpeechRequest::builder("tts-1", "hi", "alloy").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn list_voices_returns_unsupported() { + let adapter = make_adapter_no_server(); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ start / close ============ + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } +} From a8ca9d008c0d65e50185453b704aab568754134c Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 19:46:44 +0800 Subject: [PATCH 15/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1.3=20volcengine=5Fcv=20=E9=80=82=E9=85=8D=E5=99=A8=EF=BC=88?= =?UTF-8?q?=E7=81=AB=E5=B1=B1=E5=BC=95=E6=93=8E=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/adapters/volcengine_cv.rs | 1781 +++++++++++++++++ 1 file changed, 1781 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/volcengine_cv.rs diff --git a/crates/aibridge-core/src/adapters/volcengine_cv.rs b/crates/aibridge-core/src/adapters/volcengine_cv.rs new file mode 100644 index 0000000..bbf528f --- /dev/null +++ b/crates/aibridge-core/src/adapters/volcengine_cv.rs @@ -0,0 +1,1781 @@ +//! 火山引擎方舟(Seedream/Seedance)CV 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/volcengine_cv.py`。 +//! +//! 火山引擎方舟 CV 为**独立协议**(非 OpenAI 兼容),不复用 `OpenAiCompatAdapter`: +//! - 图像生成 (Seedream):`POST /images/generations`(同步) +//! - 视频生成 (Seedance):`POST /contents/generations/tasks`(异步任务,body 直传) +//! - 查询视频任务:`GET /contents/generations/tasks/{task_id}` +//! - 模型列表:`GET /models`(OpenAI 兼容端点,返回已开通模型) +//! - 认证:`Bearer ` +//! +//! 最近 3 个 Python commit 的关键对齐点(本实现已严格遵循): +//! 1. **模型 ID + 视频端点 + 请求体结构**:使用方舟 Model ID(如 `doubao-seedream-4-0-250828`), +//! 视频走 `/contents/generations/tasks`,请求体为 `{"model", "content": [...]}` 结构。 +//! 2. **图像 size 归一化到 Seedream 规范**:方舟 Seedream 对 size 有平台特定规范 +//! (最小 3686400 像素),小尺寸(如 1024x1024)会触发 503,需归一化到合法档。 +//! 3. **视频创建改用方舟官方推荐的 body 直传方式**:参数(duration/ratio/resolution/seed/ +//! watermark/camera_fixed 等)直接放 request body 顶层(强校验),不再塞进 `extra`。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::common::{infer_model_type, ModelInfo, ModelType, TaskStatus}; +use crate::model::image::{FileInput, ImageRequest, ImageResult}; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +/// 方舟默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`。 +pub const DEFAULT_VOLCENGINE_BASE_URL: &str = "https://ark.cn-beijing.volces.com/api/v3"; + +// ==================== 方舟 Seedream 图像 size 规范 ==================== +// +// 官方文档:https://www.volcengine.com/docs/82379/1541523 +// - 方式 1(枚举):"2K" / "3K" / "4K" +// - 方式 2(像素值 WIDTHxHEIGHT):需同时满足 +// * 总像素 ∈ [MIN_PIXELS, MAX_PIXELS] +// * 宽高比 ∈ [MIN_RATIO, MAX_RATIO] + +/// 方舟 Seedream 最小总像素(2560x1440 = 3686400) +const MIN_PIXELS: u64 = 3_686_400; +/// 方舟 Seedream 最大总像素(4096x4096 = 16777216) +const MAX_PIXELS: u64 = 16_777_216; +/// 方舟 Seedream 最小宽高比 +const MIN_RATIO: f64 = 1.0 / 16.0; +/// 方舟 Seedream 最大宽高比 +const MAX_RATIO: f64 = 16.0; + +/// 方舟 Seedream 合法枚举值(大写匹配) +const SIZE_ENUMS: &[&str] = &["2K", "3K", "4K"]; + +/// 2K 推荐宽高像素值表(官方推荐表,最小合法档),按宽高比升序排列 +/// +/// 格式:`(宽高比 ratio, 宽, 高)`。参考方舟推荐宽高像素值表。 +/// 不合法尺寸按最接近的宽高比映射到此表的某个档。 +static PRESETS_2K: &[(f64, u32, u32)] = &[ + (9.0 / 16.0, 1600, 2848), // 9:16 + (2.0 / 3.0, 1664, 2496), // 2:3 + (3.0 / 4.0, 1728, 2304), // 3:4 + (1.0, 2048, 2048), // 1:1 + (4.0 / 3.0, 2304, 1728), // 4:3 + (3.0 / 2.0, 2496, 1664), // 3:2 + (16.0 / 9.0, 2848, 1600), // 16:9 + (21.0 / 9.0, 3136, 1344), // 21:9 +]; + +/// 火山引擎方舟 CV 适配器(Seedream 图像 / Seedance 视频) +/// +/// 持有 HTTP 客户端与 Provider 配置,实现方舟独立协议。 +/// 不复用 `OpenAiCompatAdapter`(方舟图像/视频端点与请求体结构与 OpenAI 不一致)。 +pub struct VolcengineCvAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置 + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl VolcengineCvAdapter { + /// 创建火山引擎方舟 CV 适配器 + /// + /// `config.base_url` 为空时用 `DEFAULT_VOLCENGINE_BASE_URL` 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_VOLCENGINE_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: Self::default_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_VOLCENGINE_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: Self::default_capabilities(), + } + } + + /// 默认能力集合:图像生成 + 视频生成 + fn default_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 归一化图像 size 到方舟 Seedream 规范 + /// + /// 移植自 Python v1 `_normalize_image_size`。 + /// + /// 规则: + /// - 空/非法 → 默认 `2048x2048`(1:1 的 2K 推荐档) + /// - 枚举值(2K/3K/4K)原样透传 + /// - `WIDTHxHEIGHT`(兼容 `x`/`X`/`*` 分隔符):已合法(总像素与宽高比均在范围内)原样透传; + /// 不合法则按最接近的宽高比映射到 2K 推荐档 + pub fn normalize_image_size(size: &str) -> String { + let s = size.trim(); + if s.is_empty() { + return "2048x2048".to_string(); + } + let upper = s.to_uppercase(); + + // 方式 1:枚举值原样透传 + if SIZE_ENUMS.contains(&upper.as_str()) { + return s.to_string(); + } + + // 方式 2:解析 WIDTHxHEIGHT(兼容 x / X / * 分隔符) + let normalized = upper.replace(['X', '*'], "x"); + let parts: Vec<&str> = normalized.split('x').collect(); + if parts.len() != 2 { + return "2048x2048".to_string(); + } + let (w, h) = match (parts[0].parse::(), parts[1].parse::()) { + (Ok(w), Ok(h)) => (w, h), + _ => return "2048x2048".to_string(), + }; + if w == 0 || h == 0 { + return "2048x2048".to_string(); + } + + let total = w * h; + let ratio = w as f64 / h as f64; + + // 已合法:原样透传 + if (MIN_PIXELS..=MAX_PIXELS).contains(&total) && (MIN_RATIO..=MAX_RATIO).contains(&ratio) { + return format!("{w}x{h}"); + } + + // 不合法:按最接近的宽高比映射到 2K 推荐档 + let (best_w, best_h) = PRESETS_2K + .iter() + .min_by(|a, b| { + let da = (ratio - a.0).abs(); + let db = (ratio - b.0).abs(); + da.partial_cmp(&db).unwrap_or(std::cmp::Ordering::Equal) + }) + .map(|(_, w, h)| (*w, *h)) + .unwrap_or((2048, 2048)); + format!("{best_w}x{best_h}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: volcengine_cv)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用方舟错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 Bearer 认证的 GET 请求,并用方舟错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + // ==================== 内部:请求体构造 ==================== + + /// 构造方舟图像生成请求体 + /// + /// 移植自 Python v1 `image_generate`: + /// - size 归一化到 Seedream 规范(默认 1024x1024 → 归一化到 2K 档) + /// - n / response_format / negative_prompt / seed 透传 + /// - extra 字段合并到顶层 + fn build_image_body(&self, req: &ImageRequest) -> Value { + // 用户未指定 size 时默认 1024x1024(与 Python 一致),再归一化 + let raw_size = req.size.clone().unwrap_or_else(|| "1024x1024".to_string()); + let size = Self::normalize_image_size(&raw_size); + + let mut body = json!({ + "model": req.model, + "prompt": req.prompt, + "size": size, + "n": req.n, + "response_format": req.response_format, + }); + if let Some(np) = &req.negative_prompt { + body["negative_prompt"] = json!(np); + } + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + // extra 透传(合并到顶层) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 构造方舟视频生成请求体(body 直传方式) + /// + /// 移植自 Python v1 `video_create`(对齐方舟官方推荐 body 直传): + /// - `content` 数组:文本必选,image2video 模式附加首帧 `image_url` + /// - 视频格式参数直接放 body 顶层(强校验): + /// duration / ratio(来自 aspect_ratio) / resolution / seed / + /// watermark / camera_fixed(来自 first_frame 之外的字段,见下) + /// - 高级参数:generate_audio / service_tier / priority / draft + /// - extra 透传到顶层(覆盖同名字段) + /// + /// 注意:方舟用 `ratio` 字段表示宽高比,对应统一 `VideoRequest.aspect_ratio`。 + /// `camera_fixed` 来自统一字段 `camera_motion` 为 "fixed" 时为 true,或来自 extra.camerafixed。 + fn build_video_body(&self, req: &VideoRequest) -> Value { + // content 数组:文本必选 + let mut content = vec![json!({ "type": "text", "text": req.prompt })]; + + // image2video 模式:附加首帧 image_url + // 方舟仅支持首帧(reference_images[0]),与 Python 老版一致 + let is_image2video = matches!(req.mode, crate::model::common::VideoMode::Image2Video); + if is_image2video { + if let Some(url) = req.reference_images.first().and_then(file_input_url) { + content.push(json!({ + "type": "image_url", + "image_url": { "url": url } + })); + } else if let Some(ff) = req.first_frame.as_ref().and_then(file_input_url) { + content.push(json!({ + "type": "image_url", + "image_url": { "url": ff } + })); + } + } + + let mut body = json!({ + "model": req.model, + "content": content, + }); + + // 视频格式参数(body 直传) + if let Some(d) = req.duration { + body["duration"] = json!(d); + } + if let Some(ar) = &req.aspect_ratio { + // 方舟用 ratio 字段表示宽高比 + body["ratio"] = json!(ar); + } + if let Some(r) = &req.resolution { + body["resolution"] = json!(r); + } + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + if let Some(wm) = req.watermark { + body["watermark"] = json!(wm); + } + // camera_fixed:统一字段无直接映射,从 camera_motion == "fixed" 推断,或来自 extra + if req.camera_motion.as_deref() == Some("fixed") { + body["camera_fixed"] = json!(true); + } + + // 高级参数(仅部分模型支持,统一模型无直接字段,走 extra 透传) + // extra 透传到顶层(覆盖同名字段) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + // ==================== 内部:响应解析 ==================== + + /// 解析方舟图像生成响应 → ImageResult + /// + /// 方舟图像响应结构与 OpenAI 兼容(`{data: [{url/b64_json/revised_prompt}]}`)。 + fn parse_image_result(value: &Value, fallback_model: &str) -> Result { + let id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("img")); + let created = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + let model = value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()); + let object = value + .get("object") + .and_then(|v| v.as_str()) + .unwrap_or("image.generation") + .to_string(); + let data = value + .get("data") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_image_data_item).collect()) + .unwrap_or_default(); + Ok(ImageResult { + id, + object, + created, + model, + data, + }) + } + + /// 解析方舟视频任务创建响应 → VideoTask + /// + /// 方舟创建任务响应:`{"id": "cgt-xxx", "status": "queued", ...}` + fn parse_video_task(value: &Value, model: &str) -> Result { + let task_id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("vtask")); + let raw_status = value + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("queued"); + Ok(VideoTask { + task_id, + model: model.to_string(), + status: map_video_status(raw_status), + created_at: util::current_timestamp(), + }) + } + + /// 解析方舟视频任务查询响应 → VideoStatus + /// + /// 移植自 Python v1 `_parse_video_status`: + /// - 视频成功时 URL 在 `content.video_url`(方舟新协议),兼容 `video_url` / + /// `output.video_url` / `url` 旧路径 + /// - 失败时错误信息在 `error.message` / `message` / `error` + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let raw_status = value.get("status").and_then(|v| v.as_str()).unwrap_or(""); + let status = map_video_status(raw_status); + + // 视频 URL:方舟 content.video_url 优先,兼容其他路径 + let video_url = if status == TaskStatus::Success { + value + .get("content") + .and_then(|c| c.get("video_url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("video_url") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("video_url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| value.get("url").and_then(|v| v.as_str()).map(str::to_owned)) + } else { + None + }; + + // 错误信息:error.message / message / error + let error = if status == TaskStatus::Failed { + value + .get("error") + .and_then(|e| e.get("message")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("message") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + } else { + None + }; + + let progress = value + .get("progress") + .and_then(|v| v.as_u64()) + .map(|p| p as u32); + let created_at = value.get("created").and_then(|v| v.as_u64()); + let updated_at = value + .get("updated") + .and_then(|v| v.as_u64()) + .or_else(|| Some(util::current_timestamp())); + + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress, + error, + created_at, + updated_at, + } + } + + /// 解析方舟 /models 响应 → Vec + /// + /// 方舟 /models 为 OpenAI 兼容端点:`{"data": [{"id": "...", "created": ..., "owned_by": "..."}]}`。 + /// 返回的模型 ID 为方舟规范格式(如 `doubao-seedream-4-0-250828`)。 + fn parse_models(value: &Value) -> Vec { + match value.get("data").and_then(|v| v.as_array()) { + Some(arr) => arr + .iter() + .map(|m| { + let id = m + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let model_type = infer_model_type(&id); + ModelInfo { + name: id.clone(), + id, + model_type, + provider: "volcengine_cv".to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: None, + created: m.get("created").and_then(|v| v.as_u64()), + } + }) + .collect(), + None => Vec::new(), + } + } + + // ==================== 内部:错误映射 ==================== + + /// 将方舟 API 错误响应映射为 AibridgeError + /// + /// 移植自 Python v1 `_handle_error`: + /// - 401 → Authentication("Invalid Volcengine API key or access denied") + /// - 429 → RateLimit + /// - 404 → Api(模型端点未开通,提示检查方舟控制台 endpoint ID) + /// - 其余 ≥400 → Api(提取 error.message / message / error,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Volcengine API key or access denied".to_string(), + }, + 403 => AibridgeError::Authentication { + message: "Volcengine access denied".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Volcengine rate limit exceeded".to_string(), + retry_after: None, + }, + 404 => AibridgeError::Api { + status, + message: + "Model endpoint not found. Check your endpoint ID in Volcengine Ark console." + .to_string(), + }, + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for VolcengineCvAdapter { + fn provider_type(&self) -> &str { + "volcengine_cv" + } + + fn provider_name(&self) -> &str { + "火山引擎 CV" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 图像生成(Seedream 文生图) + /// + /// `POST /images/generations`,size 归一化到 Seedream 规范。 + async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let body = self.build_image_body(&req); + let value = self.post_authed_json("images/generations", &body).await?; + Self::parse_image_result(&value, &req.model) + } + + /// 创建视频生成任务(Seedance,body 直传方式) + /// + /// `POST /contents/generations/tasks`,参数直接放 request body 顶层(强校验)。 + async fn video_create(&self, req: VideoRequest) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let body = self.build_video_body(&req); + let value = self + .post_authed_json("contents/generations/tasks", &body) + .await?; + Self::parse_video_task(&value, &req.model) + } + + /// 查询视频任务状态 + /// + /// `GET /contents/generations/tasks/{task_id}`。 + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let path = format!("contents/generations/tasks/{task_id}"); + let value = self.get_authed_json(&path).await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 模型列表(实时拉取方舟已开通模型) + /// + /// `GET /models`,按 `filter` 过滤模型类型。 + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("models").await?; + let models = Self::parse_models(&value); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 辅助函数 ==================== + +/// 从 FileInput 提取 URL 字符串 +/// +/// - `Url(s)` → `Some(s)` +/// - `Base64(s)` → `Some(s)`(方舟接受 base64 字符串作为 image_url) +/// - `Path`/`Bytes` → `None`(方舟仅接受 URL/base64,本地路径与字节需调用方先上传) +fn file_input_url(f: &FileInput) -> Option { + match f { + FileInput::Url(s) | FileInput::Base64(s) => Some(s.clone()), + FileInput::Path(_) | FileInput::Bytes(_) => None, + } +} + +/// 映射方舟视频状态字符串到统一 TaskStatus +/// +/// 移植自 Python v1 `_map_video_status`。 +fn map_video_status(raw: &str) -> TaskStatus { + match raw.to_lowercase().as_str() { + "queued" | "pending" | "submitted" => TaskStatus::Pending, + "processing" | "running" | "in_progress" => TaskStatus::Processing, + "succeeded" | "success" | "completed" => TaskStatus::Success, + "failed" | "error" | "cancelled" => TaskStatus::Failed, + _ => TaskStatus::Pending, + } +} + +/// 解析单个 ImageData(方舟图像响应项) +/// +/// 字段与 OpenAI 兼容:url / b64_json / revised_prompt。 +fn parse_image_data_item(v: &Value) -> crate::model::image::ImageData { + crate::model::image::ImageData { + url: v.get("url").and_then(|x| x.as_str()).map(str::to_owned), + b64_json: v + .get("b64_json") + .and_then(|x| x.as_str()) + .map(str::to_owned), + revised_prompt: v + .get("revised_prompt") + .and_then(|x| x.as_str()) + .map(str::to_owned), + } +} + +/// 解析方舟错误体中的 message 字段 +/// +/// 方舟错误格式:`{"error": {"message": "..."}}` 或顶层 `{"message": "..."}`。 +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::model::common::VideoMode; + use mockito::Server; + use serde_json::json; + + /// 构造测试用 VolcengineCvAdapter(指向 mockito server) + fn make_adapter(server: &Server) -> VolcengineCvAdapter { + let opts = ClientOptions::builder() + .api_key("test-ark-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("volcengine_cv", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + VolcengineCvAdapter::with_http(http, config) + } + + /// 全能力集合(image + video) + fn full_caps() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps + } + + // ============ normalize_image_size 单元测试 ============ + + #[test] + fn normalize_size_empty_returns_default() { + assert_eq!(VolcengineCvAdapter::normalize_image_size(""), "2048x2048"); + assert_eq!( + VolcengineCvAdapter::normalize_image_size(" "), + "2048x2048" + ); + } + + #[test] + fn normalize_size_enum_passthrough() { + assert_eq!(VolcengineCvAdapter::normalize_image_size("2K"), "2K"); + assert_eq!(VolcengineCvAdapter::normalize_image_size("3K"), "3K"); + assert_eq!(VolcengineCvAdapter::normalize_image_size("4K"), "4K"); + // 小写也接受(匹配时大写化) + assert_eq!(VolcengineCvAdapter::normalize_image_size("2k"), "2k"); + } + + #[test] + fn normalize_size_valid_pixels_passthrough() { + // 2560x1440 = 3686400,正好最小像素,1:16/16:1 范围内 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("2560x1440"), + "2560x1440" + ); + // 2048x2048 = 4194304,合法 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("2048x2048"), + "2048x2048" + ); + } + + #[test] + fn normalize_size_small_1024_maps_to_2k_preset() { + // 1024x1024 = 1048576 < MIN_PIXELS,不合法,按最接近宽高比映射 + // 1:1 对应 2048x2048 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("1024x1024"), + "2048x2048" + ); + } + + #[test] + fn normalize_size_small_16x9_maps_to_2848x1600() { + // 1280x720 = 921600 < MIN_PIXELS,16:9 对应 2848x1600 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("1280x720"), + "2848x1600" + ); + } + + #[test] + fn normalize_size_small_9x16_maps_to_1600x2848() { + // 720x1280,9:16 对应 1600x2848 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("720x1280"), + "1600x2848" + ); + } + + #[test] + fn normalize_size_accepts_uppercase_x_and_star() { + // 1024X1024 与 1024*1024 等价于 1024x1024 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("2048X2048"), + "2048x2048" + ); + assert_eq!( + VolcengineCvAdapter::normalize_image_size("2048*2048"), + "2048x2048" + ); + } + + #[test] + fn normalize_size_invalid_returns_default() { + assert_eq!( + VolcengineCvAdapter::normalize_image_size("abc"), + "2048x2048" + ); + assert_eq!( + VolcengineCvAdapter::normalize_image_size("1024"), + "2048x2048" + ); + assert_eq!( + VolcengineCvAdapter::normalize_image_size("1024x"), + "2048x2048" + ); + assert_eq!( + VolcengineCvAdapter::normalize_image_size("0x1024"), + "2048x2048" + ); + } + + #[test] + fn normalize_size_too_large_maps_to_preset() { + // 8192x8192 = 67108864 > MAX_PIXELS,不合法;1:1 → 2048x2048 + assert_eq!( + VolcengineCvAdapter::normalize_image_size("8192x8192"), + "2048x2048" + ); + } + + // ============ image_generate 正常路径 ============ + + #[tokio::test] + async fn image_generate_success_parses_url() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{ + "url": "https://ark.example.com/img.png", + "revised_prompt": "a cute cat" + }] + }); + let mock = server + .mock("POST", "/images/generations") + .match_header("authorization", "Bearer test-ark-key") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "doubao-seedream-4-0-250828", + "prompt": "a cat", + "size": "2048x2048", + "n": 1, + "response_format": "url" + }))) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "a cat") + .size("1024x1024") // 小尺寸,应被归一化到 2048x2048 + .build(); + let resp = adapter.image_generate(req).await.expect("image 应成功"); + + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://ark.example.com/img.png") + ); + assert_eq!(resp.data[0].revised_prompt.as_deref(), Some("a cute cat")); + assert_eq!(resp.model, "doubao-seedream-4-0-250828"); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_size_normalization_applied() { + // 验证:用户传 1024x1024,实际发送 2048x2048(归一化生效) + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/images/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "size": "2048x2048" + }))) + .with_status(200) + .with_body( + json!({ + "created": 1, + "data": [{"url": "https://x/img.png"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat") + .size("1024x1024") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_enum_size_passthrough() { + // 2K 枚举值原样透传 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/images/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "size": "2K" + }))) + .with_status(200) + .with_body(json!({"created": 1, "data": [{"url": "https://x"}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat") + .size("2K") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_b64_response() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(200) + .with_body( + json!({ + "created": 1700000000, + "data": [{"b64_json": "aGVsbG8="}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + } + + #[tokio::test] + async fn image_generate_passes_negative_prompt_and_seed() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/images/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "negative_prompt": "blurry", + "seed": 42 + }))) + .with_status(200) + .with_body(json!({"created": 1, "data": [{"url": "https://x"}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat") + .negative_prompt("blurry") + .seed(42) + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_passes_extra_params() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/images/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "guidance_scale": 7.5 + }))) + .with_status(200) + .with_body(json!({"created": 1, "data": [{"url": "https://x"}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat") + .extra("guidance_scale", 7.5) + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + // ============ image_generate 错误路径 ============ + + #[tokio::test] + async fn image_generate_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(401) + .with_body(json!({"error": {"message": "invalid key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn image_generate_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn image_generate_error_404_returns_api_with_hint() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(404) + .with_body(json!({"error": {"message": "not found"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 404); + assert!(message.contains("Ark console")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn image_generate_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert!(message.contains("internal")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn image_generate_error_no_json_body_falls_back() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(502) + .with_body("Bad Gateway") + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("doubao-seedream-4-0-250828", "cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ video_create body 直传 正常路径 ============ + + #[tokio::test] + async fn video_create_text2video_body_direct() { + // 验证 body 直传结构:model + content[text],参数在顶层 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/contents/generations/tasks") + .match_header("authorization", "Bearer test-ark-key") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "doubao-seedance-1-0-pro-250528", + "content": [{"type": "text", "text": "a cat walking"}], + "duration": 5, + "ratio": "16:9", + "resolution": "1080p", + "seed": 123 + }))) + .with_status(200) + .with_body( + json!({ + "id": "cgt-abc123", + "status": "queued" + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "a cat walking") + .duration(5) + .aspect_ratio("16:9") + .resolution("1080p") + .seed(123) + .build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + + assert_eq!(task.task_id, "cgt-abc123"); + assert_eq!(task.model, "doubao-seedance-1-0-pro-250528"); + assert_eq!(task.status, TaskStatus::Pending); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_image2video_with_first_frame_url() { + // image2video 模式:content 含 text + image_url + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/contents/generations/tasks") + .match_body(mockito::Matcher::PartialJson(json!({ + "content": [ + {"type": "text", "text": "animate this"}, + {"type": "image_url", "image_url": {"url": "https://example.com/a.png"}} + ] + }))) + .with_status(200) + .with_body(json!({"id": "cgt-1", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "animate this") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_image2video_with_base64_first_frame() { + // base64 形式的首帧也能作为 image_url + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/contents/generations/tasks") + .match_body(mockito::Matcher::PartialJson(json!({ + "content": [ + {"type": "text", "text": "go"}, + {"type": "image_url", "image_url": {"url": "data:image/png;base64,xxx"}} + ] + }))) + .with_status(200) + .with_body(json!({"id": "cgt-2", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "go") + .mode(VideoMode::Image2Video) + .first_frame(FileInput::base64("data:image/png;base64,xxx")) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_watermark_and_camera_fixed() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/contents/generations/tasks") + .match_body(mockito::Matcher::PartialJson(json!({ + "watermark": false, + "camera_fixed": true + }))) + .with_status(200) + .with_body(json!({"id": "cgt-3", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "cat") + .watermark(false) + .camera_motion("fixed") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_extra_params_passthrough() { + // generate_audio / service_tier / priority / draft 走 extra 透传 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/contents/generations/tasks") + .match_body(mockito::Matcher::PartialJson(json!({ + "generate_audio": true, + "service_tier": "flex", + "priority": 5, + "draft": false + }))) + .with_status(200) + .with_body(json!({"id": "cgt-4", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "cat") + .extra("generate_audio", true) + .extra("service_tier", "flex") + .extra("priority", 5) + .extra("draft", false) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_uses_default_id_when_missing() { + // 方舟返回体缺 id 时,回退到生成的 vtask ID + let mut server = Server::new_async().await; + server + .mock("POST", "/contents/generations/tasks") + .with_status(200) + .with_body(json!({"status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert!(task.task_id.starts_with("vtask_")); + } + + // ============ video_create 错误路径 ============ + + #[tokio::test] + async fn video_create_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/contents/generations/tasks") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn video_create_error_429() { + let mut server = Server::new_async().await; + server + .mock("POST", "/contents/generations/tasks") + .with_status(429) + .with_body(json!({"error": {"message": "slow"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-1-0-pro-250528", "cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn video_create_error_404_hint() { + let mut server = Server::new_async().await; + server + .mock("POST", "/contents/generations/tasks") + .with_status(404) + .with_body(json!({"error": {"message": "endpoint not found"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = VideoRequest::builder("doubao-seedance-x-250528", "cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("Ark console")), + _ => panic!("应为 Api"), + } + } + + // ============ video_poll 正常路径 ============ + + #[tokio::test] + async fn video_poll_success_extracts_content_video_url() { + // 方舟新协议:视频 URL 在 content.video_url + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/contents/generations/tasks/cgt-abc") + .match_header("authorization", "Bearer test-ark-key") + .with_status(200) + .with_body( + json!({ + "id": "cgt-abc", + "status": "succeeded", + "content": {"video_url": "https://ark.example.com/v.mp4"}, + "progress": 100, + "created": 1700000000, + "updated": 1700000100 + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter + .video_poll("cgt-abc", "doubao-seedance-1-0-pro-250528") + .await + .expect("video_poll 应成功"); + + assert_eq!(status.task_id, "cgt-abc"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://ark.example.com/v.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert_eq!(status.created_at, Some(1700000000)); + assert_eq!(status.updated_at, Some(1700000100)); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_poll_processing_status() { + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-1") + .with_status(200) + .with_body( + json!({ + "id": "cgt-1", + "status": "running", + "progress": 45 + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter.video_poll("cgt-1", "m").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(45)); + assert!(status.video_url.is_none()); + } + + #[tokio::test] + async fn video_poll_failed_extracts_error_message() { + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-2") + .with_status(200) + .with_body( + json!({ + "id": "cgt-2", + "status": "failed", + "error": {"message": "content policy violation"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter.video_poll("cgt-2", "m").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("content policy violation")); + } + + #[tokio::test] + async fn video_poll_failed_extracts_top_level_message() { + // 兼容顶层 message 字段 + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-3") + .with_status(200) + .with_body( + json!({ + "id": "cgt-3", + "status": "error", + "message": "internal error" + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter.video_poll("cgt-3", "m").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("internal error")); + } + + #[tokio::test] + async fn video_poll_compatible_video_url_paths() { + // 兼容旧版 video_url / output.video_url / url 路径 + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-4") + .with_status(200) + .with_body( + json!({ + "id": "cgt-4", + "status": "completed", + "output": {"video_url": "https://legacy.example.com/v.mp4"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter.video_poll("cgt-4", "m").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://legacy.example.com/v.mp4") + ); + } + + #[tokio::test] + async fn video_poll_queued_status_maps_to_pending() { + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-5") + .with_status(200) + .with_body(json!({"id": "cgt-5", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let status = adapter.video_poll("cgt-5", "m").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn video_poll_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/contents/generations/tasks/cgt-x") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let err = adapter.video_poll("cgt-x", "m").await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .match_header("authorization", "Bearer test-ark-key") + .with_status(200) + .with_body(json!({ + "object": "list", + "data": [ + {"id": "doubao-seedream-4-0-250828", "object": "model", "created": 1, "owned_by": "volcengine"}, + {"id": "doubao-seedance-1-0-pro-250528", "object": "model", "created": 1, "owned_by": "volcengine"}, + {"id": "doubao-pro-32k", "object": "model", "created": 1, "owned_by": "volcengine"} + ] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + // 类型推断:seedream→image,seedance→video,doubao-pro→chat + assert_eq!(models[0].model_type, ModelType::Image); + assert_eq!(models[1].model_type, ModelType::Video); + assert_eq!(models[2].model_type, ModelType::Chat); + // provider 字段填充 + assert_eq!(models[0].provider, "volcengine_cv"); + } + + #[tokio::test] + async fn list_models_filter_image() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "doubao-seedream-4-0-250828", "created": 1}, + {"id": "doubao-seedance-1-0-pro-250528", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "doubao-seedream-4-0-250828"); + } + + #[tokio::test] + async fn list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_models_empty_data() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body(json!({"data": []}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert!(models.is_empty()); + } + + // ============ Adapter trait 元信息 ============ + + #[tokio::test] + async fn provider_metadata() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + assert_eq!(adapter.provider_type(), "volcengine_cv"); + assert_eq!(adapter.provider_name(), "火山引擎 CV"); + assert!(adapter.requires_api_key()); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::VideoGenerate)); + } + + #[tokio::test] + async fn unsupported_chat_returns_error() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let req = crate::model::chat::ChatRequest::builder("m", vec![]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn unsupported_embed_returns_error() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let req = crate::model::options::EmbedRequest { + model: "m".into(), + input: crate::model::options::EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: std::collections::HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut server = Server::new_async().await; + let mut adapter = make_adapter(&server); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + // 抑制未使用 mut 警告 + let _ = &mut server; + } + + // ============ base_url 选择 ============ + + #[test] + fn base_url_uses_config_when_provided() { + let config = ProviderConfig::from_options( + "volcengine_cv", + ClientOptions::builder() + .api_key("k") + .base_url("https://custom.ark.example.com/api/v3") + .build(), + ); + let adapter = VolcengineCvAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.ark.example.com/api/v3"); + } + + #[test] + fn base_url_falls_back_to_default_when_missing() { + let config = ProviderConfig::from_options( + "volcengine_cv", + ClientOptions::builder().api_key("k").build(), + ); + let adapter = VolcengineCvAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_VOLCENGINE_BASE_URL); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn map_api_error_401_is_authentication() { + let err = VolcengineCvAdapter::map_api_error(401, "{\"error\":{\"message\":\"bad\"}}"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_api_error_429_is_rate_limit() { + let err = VolcengineCvAdapter::map_api_error(429, "{}"); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn map_api_error_404_has_console_hint() { + let err = VolcengineCvAdapter::map_api_error(404, "{}"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("Ark console")), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_500_extracts_message() { + let err = VolcengineCvAdapter::map_api_error(500, "{\"error\":{\"message\":\"internal\"}}"); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_top_level_message() { + // 顶层 message 字段 + let err = VolcengineCvAdapter::map_api_error(400, "{\"message\":\"bad param\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "bad param"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_falls_back_to_http_status() { + let err = VolcengineCvAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ 状态映射单元测试 ============ + + #[test] + fn map_video_status_variants() { + assert_eq!(map_video_status("queued"), TaskStatus::Pending); + assert_eq!(map_video_status("PENDING"), TaskStatus::Pending); + assert_eq!(map_video_status("submitted"), TaskStatus::Pending); + assert_eq!(map_video_status("processing"), TaskStatus::Processing); + assert_eq!(map_video_status("running"), TaskStatus::Processing); + assert_eq!(map_video_status("in_progress"), TaskStatus::Processing); + assert_eq!(map_video_status("succeeded"), TaskStatus::Success); + assert_eq!(map_video_status("success"), TaskStatus::Success); + assert_eq!(map_video_status("completed"), TaskStatus::Success); + assert_eq!(map_video_status("failed"), TaskStatus::Failed); + assert_eq!(map_video_status("error"), TaskStatus::Failed); + assert_eq!(map_video_status("cancelled"), TaskStatus::Failed); + // 未知状态默认 Pending + assert_eq!(map_video_status("unknown"), TaskStatus::Pending); + assert_eq!(map_video_status(""), TaskStatus::Pending); + } + + // ============ file_input_url 辅助 ============ + + #[test] + fn file_input_url_extracts_url() { + assert_eq!( + file_input_url(&FileInput::url("https://x")), + Some("https://x".to_string()) + ); + } + + #[test] + fn file_input_url_extracts_base64() { + assert_eq!( + file_input_url(&FileInput::base64("aGk=")), + Some("aGk=".to_string()) + ); + } + + #[test] + fn file_input_url_rejects_path_and_bytes() { + assert_eq!(file_input_url(&FileInput::path("/tmp/x")), None); + assert_eq!(file_input_url(&FileInput::bytes(vec![1, 2])), None); + } + + // ============ full_caps 占位(保持与 openai_compat 测试结构一致)============ + + #[test] + fn full_caps_contains_image_and_video() { + let caps = full_caps(); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::VideoGenerate)); + } +} From cc2026e36808b61480af60e895f6d0cf79750d98 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 19:53:13 +0800 Subject: [PATCH 16/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1.4=20gemini=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/gemini.rs | 2264 +++++++++++++++++++ 1 file changed, 2264 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/gemini.rs diff --git a/crates/aibridge-core/src/adapters/gemini.rs b/crates/aibridge-core/src/adapters/gemini.rs new file mode 100644 index 0000000..55bf1af --- /dev/null +++ b/crates/aibridge-core/src/adapters/gemini.rs @@ -0,0 +1,2264 @@ +//! Google Gemini 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/gemini.py`。 +//! +//! Gemini 是独立协议(非 OpenAI 兼容),核心端点: +//! - 对话:`POST /models/{model}:generateContent` +//! - 流式对话:`POST /models/{model}:streamGenerateContent?alt=sse` +//! - 嵌入:`POST /models/{model}:embedContent`(单条)/ `:batchEmbedContents`(多条) +//! - 模型列表:`GET /models` +//! - 认证:API Key 通过 `x-goog-api-key` header 传递(也可作 `key` query param) +//! +//! 设计要点(与设计文档第 10 节一致): +//! - 独立协议,不复用 `OpenAiCompatAdapter` +//! - 请求体用 Gemini 的 `contents`/`parts` 结构,system 消息转 `systemInstruction` +//! - assistant 角色 → Gemini 的 `model` 角色 +//! - 多模态 `image_url` (data URI) → Gemini 的 `inline_data` (mime_type + data) +//! - 参数映射 `GEMINI_MAPPING`:max_tokens→maxOutputTokens、top_p→topP、top_k→topK、stop→stopSequences +//! - 错误映射:Gemini 错误体 `{"error":{"code","message","status"}}` → AibridgeError +//! - 图像生成:通过 `:generateContent` + `generationConfig.responseModalities:["IMAGE","TEXT"]` +//! 调用 Gemini 图像生成模型(如 `gemini-2.0-flash-exp-image-generation`), +//! 响应 `parts` 中的 `inlineData.data` (base64) 转为统一 `ImageData.b64_json` + +use std::collections::HashMap; + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChoiceMessage, DeltaMessage, UserContent, +}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::image::{ImageData, ImageRequest, ImageResult}; +use crate::model::options::{ + EmbedInput, EmbedRequest, EmbeddingItem, EmbeddingResult, EmbeddingVector, ParameterMapping, +}; +use crate::util; + +/// Gemini 官方默认 Base URL +pub const DEFAULT_GEMINI_BASE_URL: &str = "https://generativelanguage.googleapis.com/v1beta"; + +/// Gemini API key 请求头名 +const API_KEY_HEADER: &str = "x-goog-api-key"; + +/// Gemini 通用参数映射 +/// +/// 对应 Python v1 `agn/models/options.py` 的 `GEMINI_MAPPING`。 +/// Gemini 的 generationConfig 字段命名与统一参数不同,需重命名: +/// - `max_tokens` → `maxOutputTokens` +/// - `top_p` → `topP` +/// - `top_k` → `topK` +/// - `stop` → `stopSequences` +/// - `response_format` → `response_mime_type`(值需额外转换,此处仅重命名键) +/// +/// 注:Python 老版的 `value_map`(reasoning/web_search)在 Rust 统一请求中 +/// 走 `extra` 透传更合适,故此处仅保留 rename_map。 +pub fn gemini_mapping() -> ParameterMapping { + let mut rename_map = HashMap::new(); + rename_map.insert( + "max_tokens".to_string(), + Some("maxOutputTokens".to_string()), + ); + rename_map.insert("top_p".to_string(), Some("topP".to_string())); + rename_map.insert("top_k".to_string(), Some("topK".to_string())); + rename_map.insert("stop".to_string(), Some("stopSequences".to_string())); + rename_map.insert( + "response_format".to_string(), + Some("response_mime_type".to_string()), + ); + ParameterMapping { rename_map } +} + +/// Google Gemini 适配器 +/// +/// 持有 HTTP 客户端与 Provider 配置,实现 Gemini 独立协议。 +/// 不复用 `OpenAiCompatAdapter`(Gemini 请求/响应结构、认证方式、错误格式均不同)。 +pub struct GeminiAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl GeminiAdapter { + /// 创建 Gemini 适配器 + /// + /// `config.base_url` 为空时用 `DEFAULT_GEMINI_BASE_URL` 兜底。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_GEMINI_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: Self::default_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_GEMINI_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: Self::default_capabilities(), + } + } + + /// 默认能力集合:CHAT / IMAGE_GENERATE / EMBEDDING + fn default_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::Vision); + caps + } + + /// API key(可能为空) + fn api_key(&self) -> Option<&str> { + self.config.api_key.as_deref().filter(|s| !s.is_empty()) + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: gemini)", cap.as_str()), + }) + } + } + + /// 校验 API key 是否存在(requires_api_key=true) + fn ensure_api_key(&self) -> Result<&str> { + match self.api_key() { + Some(k) => Ok(k), + None => Err(AibridgeError::Validation { + message: "Gemini 适配器需要 API key(x-goog-api-key)".into(), + details: serde_json::json!({ "provider": "gemini" }), + }), + } + } + + // ==================== 消息格式转换 ==================== + + /// 将统一 `ChatRequest` 转为 Gemini 请求体 + /// + /// - system 消息 → `systemInstruction`(单独字段,不在 contents 中) + /// - user 消息 → `contents` 中 role=user + /// - assistant 消息 → `contents` 中 role=model(Gemini 用 model 表示助手) + /// - tool 消息 → 作为 user 内容追加(Gemini 无独立 tool role) + /// - 多模态 image_url (data URI) → `inline_data` (mime_type + data) + /// - 通用参数(temperature/top_p/top_k/max_tokens/stop)→ `generationConfig`,经 GEMINI_MAPPING 重命名 + /// - extra 透传到顶层 + fn build_chat_body(&self, req: &ChatRequest) -> Value { + let mut contents: Vec = Vec::new(); + let mut system_instruction: Option = None; + + for msg in &req.messages { + match msg { + ChatMessage::System { content, .. } => { + if system_instruction.is_none() { + system_instruction = Some(content.clone()); + } + } + ChatMessage::User { content, .. } => { + let parts = Self::convert_user_content(content); + contents.push(json!({ "role": "user", "parts": parts })); + } + ChatMessage::Assistant { content, .. } => { + let text = content.clone().unwrap_or_default(); + contents.push(json!({ "role": "model", "parts": [{ "text": text }] })); + } + ChatMessage::Tool { content, .. } => { + // 工具结果作为 user 消息内容追加(Gemini 无 tool role) + contents.push(json!({ "role": "user", "parts": [{ "text": content }] })); + } + } + } + + let mut body = json!({ "contents": contents }); + + if let Some(sys) = system_instruction { + body["systemInstruction"] = json!({ "parts": [{ "text": sys }] }); + } + + // generationConfig:经 GEMINI_MAPPING 重命名通用参数 + let gen_config = self.build_generation_config(req); + if !gen_config.is_null() { + body["generationConfig"] = gen_config; + } + + // extra 透传到顶层 + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + + body + } + + /// 构造 generationConfig(应用 GEMINI_MAPPING 重命名) + /// + /// 只收集非空的通用生成参数。返回 `Value::Null` 表示无参数。 + fn build_generation_config(&self, req: &ChatRequest) -> Value { + let mut params: HashMap = HashMap::new(); + if let Some(t) = req.temperature { + params.insert("temperature".into(), json!(t)); + } + if let Some(t) = req.top_p { + params.insert("top_p".into(), json!(t)); + } + if let Some(t) = req.top_k { + params.insert("top_k".into(), json!(t)); + } + if let Some(t) = req.max_tokens { + params.insert("max_tokens".into(), json!(t)); + } + if let Some(stop) = &req.stop { + // StopSeq 可能是 Single 或 Multiple,统一转数组 + let arr: Vec = match stop { + crate::model::options::StopSeq::Single(s) => vec![s.clone()], + crate::model::options::StopSeq::Multiple(v) => v.clone(), + }; + params.insert("stop".into(), json!(arr)); + } + + let mapped = gemini_mapping().apply(¶ms); + if mapped.is_empty() { + return Value::Null; + } + let mut obj = serde_json::Map::new(); + for (k, v) in mapped { + obj.insert(k, v); + } + Value::Object(obj) + } + + /// 将统一 UserContent 转为 Gemini parts 数组 + /// + /// - 纯文本 → `[{"text": "..."}]` + /// - 多模态 → 文本部件 + 图像 inline_data(data URI 解析出 mime_type + base64 data) + fn convert_user_content(content: &UserContent) -> Vec { + match content { + UserContent::Text(s) => vec![json!({ "text": s })], + UserContent::Parts(parts) => parts + .iter() + .filter_map(|p| match p { + crate::model::chat::ContentPart::Text { text } => Some(json!({ "text": text })), + crate::model::chat::ContentPart::ImageUrl { image_url } => { + Self::convert_image_url(&image_url.url) + } + }) + .collect(), + } + } + + /// 将 OpenAI 风格 image_url(URL 或 data URI)转为 Gemini part + /// + /// - data URI (`data:image/png;base64,xxx`) → `inline_data` (mime_type + data) + /// - 普通 URL → `file_data` (file_uri)(Gemini 用 file_data 引用远程文件) + fn convert_image_url(url: &str) -> Option { + if let Some(rest) = url.strip_prefix("data:") { + // data URI 格式: data:;base64, + if let Some((meta, data)) = rest.split_once(',') { + let mime_type = meta.split(';').next().unwrap_or("image/png"); + return Some(json!({ + "inline_data": { "mime_type": mime_type, "data": data } + })); + } + } + // 普通远程 URL:用 file_data 引用 + Some(json!({ + "file_data": { "file_uri": url, "mime_type": "image/png" } + })) + } + + // ==================== HTTP 请求 ==================== + + /// 发送带认证的 POST JSON 请求,用 Gemini 错误映射处理响应 + async fn post_authed_json(&self, url: &str, body: &Value) -> Result { + let api_key = self.ensure_api_key()?; + let resp = self + .http + .inner() + .post(url) + .header(API_KEY_HEADER, api_key) + .json(body) + .send() + .await + .map_err(map_reqwest_error)?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带认证的 GET 请求,用 Gemini 错误映射处理响应 + async fn get_authed_json(&self, url: &str) -> Result { + let api_key = self.ensure_api_key()?; + let resp = self + .http + .inner() + .get(url) + .header(API_KEY_HEADER, api_key) + .send() + .await + .map_err(map_reqwest_error)?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + // ==================== 响应解析 ==================== + + /// 解析 Gemini generateContent 响应 → ChatCompletion + /// + /// Gemini 响应格式: + /// ```json + /// { + /// "candidates": [{ + /// "content": { "parts": [{"text":"..."}], "role": "model" }, + /// "finishReason": "STOP" + /// }], + /// "usageMetadata": { + /// "promptTokenCount": 5, + /// "candidatesTokenCount": 10, + /// "totalTokenCount": 15 + /// } + /// } + /// ``` + fn parse_chat_completion(&self, value: &Value, fallback_model: &str) -> Result { + let candidates = value.get("candidates").and_then(|v| v.as_array()); + let (content, finish_reason) = match candidates.and_then(|arr| arr.first()) { + Some(cand) => { + let parts = cand + .get("content") + .and_then(|c| c.get("parts")) + .and_then(|p| p.as_array()); + let text = parts + .map(|arr| { + arr.iter() + .filter_map(|p| p.get("text").and_then(|t| t.as_str())) + .collect::>() + .join("") + }) + .unwrap_or_default(); + let raw_reason = cand + .get("finishReason") + .and_then(|v| v.as_str()) + .unwrap_or("STOP"); + (text, map_finish_reason(raw_reason)) + } + None => (String::new(), "stop".to_string()), + }; + + let usage = value.get("usageMetadata").and_then(parse_usage_metadata); + + Ok(ChatCompletion { + id: util::generate_id("chatcmpl"), + object: "chat.completion".into(), + created: util::current_timestamp(), + model: fallback_model.to_string(), + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".into(), + content: Some(content), + tool_calls: None, + }, + finish_reason: Some(finish_reason), + }], + usage, + service_tier: None, + system_fingerprint: None, + }) + } + + /// 解析单个 SSE chunk(Gemini streamGenerateContent 格式)→ Option + /// + /// Gemini 每个 SSE data 行是一个完整的 generateContent 响应(含 candidates)。 + /// 文本增量取 candidates[0].content.parts 的 text 拼接;finishReason 出现时单独发一个结束块。 + /// 返回 None 表示该 chunk 无有效内容(如空 candidates),调用方跳过。 + fn parse_stream_chunk(value: &Value, model: &str) -> Result> { + let id = util::generate_id("chatcmpl"); + let created = util::current_timestamp(); + let candidates = match value.get("candidates").and_then(|v| v.as_array()) { + Some(arr) => arr, + None => return Ok(Vec::new()), + }; + + let mut chunks = Vec::new(); + if let Some(cand) = candidates.first() { + let parts = cand + .get("content") + .and_then(|c| c.get("parts")) + .and_then(|p| p.as_array()); + let text = parts + .map(|arr| { + arr.iter() + .filter_map(|p| p.get("text").and_then(|t| t.as_str())) + .collect::>() + .join("") + }) + .unwrap_or_default(); + + if !text.is_empty() { + chunks.push(ChatCompletionChunk { + id: id.clone(), + object: "chat.completion.chunk".into(), + created, + model: model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: None, + content: Some(text), + tool_calls: None, + }, + finish_reason: None, + }], + usage: None, + }); + } + + // finishReason 出现时发结束块 + if let Some(raw) = cand.get("finishReason").and_then(|v| v.as_str()) { + chunks.push(ChatCompletionChunk { + id, + object: "chat.completion.chunk".into(), + created, + model: model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: None, + content: None, + tool_calls: None, + }, + finish_reason: Some(map_finish_reason(raw)), + }], + usage: None, + }); + } + } + + // 末尾可能携带 usageMetadata + if let Some(usage) = value.get("usageMetadata").and_then(parse_usage_metadata) { + chunks.push(ChatCompletionChunk { + id: util::generate_id("chatcmpl"), + object: "chat.completion.chunk".into(), + created, + model: model.to_string(), + choices: Vec::new(), + usage: Some(usage), + }); + } + + Ok(chunks) + } + + /// 解析 Gemini /models 响应 → Vec + /// + /// Gemini 响应特殊:模型列表在 `models` 键下(非 OpenAI 的 `data`), + /// 模型 ID 在 `name` 字段且带 `models/` 前缀(如 `models/gemini-1.5-pro`), + /// 显示名在 `displayName` 字段。需去掉前缀作为统一 id。 + fn parse_models(&self, value: &Value) -> Vec { + let arr = value.get("models").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .filter_map(|m| { + let name = m.get("name").and_then(|v| v.as_str())?; + // 去掉 "models/" 前缀作为模型 ID + let id = name.strip_prefix("models/").unwrap_or(name).to_string(); + if id.is_empty() { + return None; + } + let display_name = m + .get("displayName") + .and_then(|v| v.as_str()) + .unwrap_or(&id) + .to_string(); + let model_type = infer_model_type(&id); + let supports_streaming = matches!(model_type, ModelType::Chat); + Some(ModelInfo { + name: display_name, + id, + model_type, + provider: "gemini".into(), + capabilities: Vec::new(), + max_tokens: m + .get("inputTokenLimit") + .and_then(|v| v.as_u64()) + .map(|x| x as u32), + supports_streaming, + description: m + .get("description") + .and_then(|v| v.as_str()) + .map(str::to_owned), + created: None, + }) + }) + .collect(), + None => Vec::new(), + } + } + + // ==================== 错误映射 ==================== + + /// 将 Gemini API 错误响应映射为 AibridgeError + /// + /// Gemini 错误格式: + /// ```json + /// { "error": { "code": 400, "message": "...", "status": "INVALID_ARGUMENT" } } + /// ``` + /// 状态码分类: + /// - 401/403 → Authentication + /// - 429 → RateLimit + /// - 404 → ModelNotFound + /// - 400 → Validation + /// - 4xx(其他)→ Api + /// - 5xx → Api + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + let message = parse_gemini_error_message(body, status); + match status { + 401 | 403 => AibridgeError::Authentication { message }, + 429 => AibridgeError::RateLimit { + message, + retry_after: None, + }, + 400 => AibridgeError::Validation { + message, + details: serde_json::json!({ "status_code": status, "response": body }), + }, + 404 => AibridgeError::ModelNotFound { model: message }, + s => AibridgeError::Api { status: s, message }, + } + } +} + +#[async_trait] +impl Adapter for GeminiAdapter { + fn provider_type(&self) -> &str { + "gemini" + } + + fn provider_name(&self) -> &str { + "Google Gemini" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 已在 new() 中初始化,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 通过 Drop 自动释放连接池,无需显式关闭 + Ok(()) + } + + /// 文本对话 + /// + /// `POST /models/{model}:generateContent`,构造 Gemini 请求体,解析响应 → ChatCompletion。 + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let model = req.model.clone(); + let body = self.build_chat_body(&req); + let url = self.url(&format!("models/{model}:generateContent")); + let value = self.post_authed_json(&url, &body).await?; + self.parse_chat_completion(&value, &model) + } + + /// 流式文本对话 + /// + /// `POST /models/{model}:streamGenerateContent?alt=sse`,SSE 流 → ChatStream。 + /// Gemini 流式格式:每个 `data:` 行是一个完整 generateContent 响应(含 candidates)。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let model = req.model.clone(); + let body = self.build_chat_body(&req); + let url = self.url(&format!("models/{model}:streamGenerateContent?alt=sse")); + + let api_key = self.ensure_api_key()?; + let resp = self + .http + .inner() + .post(&url) + .header(API_KEY_HEADER, api_key) + .json(&body) + .send() + .await + .map_err(map_reqwest_error)?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + + // 按字节流读取,按行切分解析 SSE(与 openai_compat 一致的 LinesStream 模式) + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = LinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 去除 "data: " 前缀 + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + // Gemini 流无 [DONE] 标记,自然结束即结束 + match serde_json::from_str::(data) { + Ok(v) => match Self::parse_stream_chunk(&v, &model) { + Ok(chunks) => { + for chunk in chunks { + yield Ok(chunk); + } + } + Err(e) => { + yield Err(e); + return; + } + }, + Err(_) => { + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + continue; + } + } + } + }; + + Ok(stream.boxed()) + } + + /// 图像生成 + /// + /// 通过 `:generateContent` 端点调用 Gemini 图像生成模型 + /// (如 `gemini-2.0-flash-exp-image-generation`)。 + /// 请求体 `generationConfig.responseModalities: ["IMAGE","TEXT"]`, + /// 响应 `candidates[0].content.parts` 中的 `inlineData.data` (base64) 转为 `ImageData.b64_json`。 + async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let model = req.model.clone(); + let prompt = req.prompt.clone(); + + let mut body = json!({ + "contents": [{ "role": "user", "parts": [{ "text": prompt }] }], + "generationConfig": { "responseModalities": ["IMAGE", "TEXT"] } + }); + + // 透传 extra(如 negative_prompt 等 provider 特有参数) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + + let url = self.url(&format!("models/{model}:generateContent")); + let value = self.post_authed_json(&url, &body).await?; + + // 提取响应中的图像(inlineData → b64_json) + let mut data: Vec = Vec::new(); + if let Some(parts) = value + .get("candidates") + .and_then(|c| c.as_array()) + .and_then(|arr| arr.first()) + .and_then(|cand| cand.get("content")) + .and_then(|c| c.get("parts")) + .and_then(|p| p.as_array()) + { + for part in parts { + if let Some(inline) = part.get("inlineData").or_else(|| part.get("inline_data")) { + let b64 = inline + .get("data") + .and_then(|d| d.as_str()) + .map(str::to_owned); + let mime = inline + .get("mimeType") + .or_else(|| inline.get("mime_type")) + .and_then(|m| m.as_str()) + .unwrap_or("image/png"); + if let Some(b64) = b64 { + data.push(ImageData { + url: None, + b64_json: Some(b64), + revised_prompt: Some(format!("Generated image ({mime})")), + }); + } + } + } + } + + Ok(ImageResult { + id: util::generate_id("img"), + object: "image.generation".into(), + created: util::current_timestamp(), + model, + data, + }) + } + + /// 文本嵌入 + /// + /// - 单条输入:`POST /models/{model}:embedContent` + /// - 多条输入:`POST /models/{model}:batchEmbedContents` + /// + /// Gemini 嵌入响应无 usage 统计,故 `usage` 为 None。 + async fn embed(&self, req: EmbedRequest) -> Result { + self.ensure_capability(Capabilities::Embedding)?; + let model = req.model.clone(); + + let result = match &req.input { + EmbedInput::Single(text) => { + let mut body = json!({ + "model": format!("models/{model}"), + "content": { "parts": [{ "text": text }] } + }); + if let Some(dim) = req.dimensions { + body["outputDimensionality"] = json!(dim); + } + let url = self.url(&format!("models/{model}:embedContent")); + let value = self.post_authed_json(&url, &body).await?; + + let values = value + .get("embedding") + .and_then(|e| e.get("values")) + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().filter_map(|x| x.as_f64()).collect::>()) + .unwrap_or_default(); + vec![EmbeddingItem { + object: "embedding".into(), + index: 0, + embedding: EmbeddingVector::Float(values), + }] + } + EmbedInput::Multiple(texts) => { + let requests: Vec = texts + .iter() + .map(|text| { + let mut r = json!({ + "model": format!("models/{model}"), + "content": { "parts": [{ "text": text }] } + }); + if let Some(dim) = req.dimensions { + r["outputDimensionality"] = json!(dim); + } + r + }) + .collect(); + let body = json!({ "requests": requests }); + let url = self.url(&format!("models/{model}:batchEmbedContents")); + let value = self.post_authed_json(&url, &body).await?; + + value + .get("embeddings") + .and_then(|e| e.as_array()) + .map(|arr| { + arr.iter() + .enumerate() + .map(|(i, emb)| { + let values = emb + .get("values") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter().filter_map(|x| x.as_f64()).collect::>() + }) + .unwrap_or_default(); + EmbeddingItem { + object: "embedding".into(), + index: i as u32, + embedding: EmbeddingVector::Float(values), + } + }) + .collect() + }) + .unwrap_or_default() + } + }; + + Ok(EmbeddingResult { + object: "list".into(), + data: result, + model, + usage: None, + }) + } + + /// 模型列表(实时拉取) + /// + /// `GET /models` → Vec,按 `filter` 过滤模型类型。 + /// Gemini 响应的 `models` 数组中 `name` 带 `models/` 前缀,需去掉。 + async fn list_models(&self, filter: Option) -> Result> { + let url = self.url("models"); + let value = self.get_authed_json(&url).await?; + let models = self.parse_models(&value); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 辅助函数 ==================== + +/// 将 reqwest::Error 映射为 AibridgeError(超时 → Timeout,其余 → Network) +fn map_reqwest_error(err: reqwest::Error) -> AibridgeError { + if err.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(err) + } +} + +/// 解析 Gemini 错误体中的 message 字段 +/// +/// Gemini 错误格式:`{"error": {"code": 400, "message": "...", "status": "INVALID_ARGUMENT"}}` +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_gemini_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 部分错误体直接用顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +/// 将 Gemini finishReason 映射为统一 finish_reason +/// +/// - STOP → stop +/// - MAX_TOKENS → length +/// - SAFETY → content_filter +/// - RECITATION → stop(Python 老版映射为 stop) +/// - 其他 → stop(兜底) +fn map_finish_reason(raw: &str) -> String { + match raw { + "STOP" => "stop".into(), + "MAX_TOKENS" => "length".into(), + "SAFETY" => "content_filter".into(), + "RECITATION" => "stop".into(), + _ => "stop".into(), + } +} + +/// 解析 Gemini usageMetadata → ChatUsage +/// +/// Gemini 字段:promptTokenCount / candidatesTokenCount / totalTokenCount +fn parse_usage_metadata(v: &Value) -> Option { + let prompt = v.get("promptTokenCount").and_then(|x| x.as_u64())?; + let completion = v + .get("candidatesTokenCount") + .and_then(|x| x.as_u64()) + .unwrap_or(0); + let total = v + .get("totalTokenCount") + .and_then(|x| x.as_u64()) + .unwrap_or(prompt + completion); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: total, + }) +} + +// ==================== SSE 行流适配器(与 openai_compat 一致) ==================== + +/// 将字节流按行切分的适配器 +/// +/// 维护一个未完成行的缓冲区,逐 chunk 拼接出完整行。 +struct LinesStream { + inner: S, + buffer: Vec, +} + +impl LinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for LinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + } + Poll::Ready(None) => { + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::chat::ContentPart; + use crate::model::image::FileInput; + use mockito::Server; + + /// 构造测试用 GeminiAdapter(指向 mockito server) + fn make_adapter(server: &Server) -> GeminiAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("gemini", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + GeminiAdapter::with_http(http, config) + } + + /// 构造无 API key 的适配器(测试 ensure_api_key) + fn make_adapter_no_key(server: &Server) -> GeminiAdapter { + let opts = ClientOptions::builder() + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("gemini", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + GeminiAdapter::with_http(http, config) + } + + // ============ chat 正常路径 ============ + + #[tokio::test] + async fn chat_success_parses_completion() { + let mut server = Server::new_async().await; + let body = json!({ + "candidates": [{ + "content": { + "parts": [{"text": "Hello from Gemini!"}], + "role": "model" + }, + "finishReason": "STOP" + }], + "usageMetadata": { + "promptTokenCount": 5, + "candidatesTokenCount": 3, + "totalTokenCount": 8 + } + }); + let mock = server + .mock("POST", "/models/gemini-2.5-pro:generateContent") + .match_header("x-goog-api-key", "test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gemini-2.5-pro", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.model, "gemini-2.5-pro"); + assert_eq!(resp.choices.len(), 1); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("Hello from Gemini!") + ); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 8); + assert_eq!(resp.usage.as_ref().unwrap().prompt_tokens, 5); + assert_eq!(resp.usage.as_ref().unwrap().completion_tokens, 3); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_converts_messages_to_gemini_format() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/gemini-2.5-pro:generateContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "contents": [ + { "role": "user", "parts": [{ "text": "hello" }] }, + { "role": "model", "parts": [{ "text": "hi there" }] }, + { "role": "user", "parts": [{ "text": "how are you?" }] } + ], + "systemInstruction": { "parts": [{ "text": "You are helpful." }] } + }))) + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{"text": "ok"}], "role": "model" }, + "finishReason": "STOP" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder( + "gemini-2.5-pro", + vec![ + ChatMessage::system("You are helpful."), + ChatMessage::user("hello"), + ChatMessage::assistant("hi there"), + ChatMessage::user("how are you?"), + ], + ) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_applies_generation_config_mapping() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/gemini-2.5-pro:generateContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "generationConfig": { + "temperature": 0.7, + "topP": 0.9, + "topK": 40, + "maxOutputTokens": 1000, + "stopSequences": ["END"] + } + }))) + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{"text": "ok"}], "role": "model" }, + "finishReason": "STOP" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gemini-2.5-pro", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .top_p(0.9) + .top_k(40) + .max_tokens(1000) + .stop(crate::model::options::StopSeq::Single("END".into())) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_passes_extra_to_top_level() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/gemini-2.5-pro:generateContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "contents": [{ "role": "user", "parts": [{ "text": "hi" }] }], + "thinkingConfig": { "thinkingBudget": -1 } + }))) + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{"text": "ok"}], "role": "model" }, + "finishReason": "STOP" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gemini-2.5-pro", vec![ChatMessage::user("hi")]) + .extra("thinkingConfig", json!({ "thinkingBudget": -1 })) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_finish_reason_mappings() { + let cases = [ + ("STOP", "stop"), + ("MAX_TOKENS", "length"), + ("SAFETY", "content_filter"), + ("RECITATION", "stop"), + ("OTHER", "stop"), + ]; + for (raw, expected) in cases { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{"text": "x"}], "role": "model" }, + "finishReason": raw + }] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!( + resp.choices[0].finish_reason.as_deref(), + Some(expected), + "finishReason={raw} 应映射为 {expected}" + ); + } + } + + #[tokio::test] + async fn chat_multimodal_converts_data_uri_to_inline_data() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/m:generateContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "contents": [{ + "role": "user", + "parts": [ + { "text": "describe this" }, + { "inline_data": { "mime_type": "image/png", "data": "aGVsbG8=" } } + ] + }] + }))) + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{"text": "a cat"}], "role": "model" }, + "finishReason": "STOP" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder( + "m", + vec![ChatMessage::user_multimodal(vec![ + ContentPart::Text { + text: "describe this".into(), + }, + ContentPart::ImageUrl { + image_url: crate::model::chat::ImageUrl::new("data:image/png;base64,aGVsbG8="), + }, + ])], + ) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_no_candidates_returns_empty_content() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(200) + .with_body(json!({}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("")); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + } + + // ============ chat 错误路径 ============ + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(401) + .with_body( + json!({"error": {"code": 401, "message": "API key not valid", "status": "UNAUTHENTICATED"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Authentication { message } => { + assert!(message.contains("API key not valid")); + } + _ => panic!("应为 Authentication,实际: {err:?}"), + } + } + + #[tokio::test] + async fn chat_error_403_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(403) + .with_body(json!({"error": {"message": "permission denied"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(429) + .with_body( + json!({"error": {"code": 429, "message": "Resource exhausted", "status": "RESOURCE_EXHAUSTED"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::RateLimit { + message, + retry_after, + } => { + assert!(message.contains("Resource exhausted")); + assert_eq!(retry_after, None); + } + _ => panic!("应为 RateLimit"), + } + } + + #[tokio::test] + async fn chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/gemini-x:generateContent") + .with_status(404) + .with_body( + json!({"error": {"code": 404, "message": "models/gemini-x not found"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gemini-x", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::ModelNotFound { model } => { + assert!(model.contains("gemini-x")); + } + _ => panic!("应为 ModelNotFound"), + } + } + + #[tokio::test] + async fn chat_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(400) + .with_body( + json!({"error": {"code": 400, "message": "Invalid argument", "status": "INVALID_ARGUMENT"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("Invalid argument")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(500) + .with_body(json!({"error": {"message": "Internal error"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn chat_missing_api_key_returns_validation() { + let server = Server::new_async().await; + let adapter = make_adapter_no_key(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn chat_unsupported_capability_returns_error() { + // GeminiAdapter 默认支持 Chat,这里通过空能力集合的 adapter 测试 + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("gemini", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let mut adapter = GeminiAdapter::with_http(http, config); + adapter.capabilities = CapabilitySet::new(); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ chat_stream 正常 + 错误路径 ============ + + #[tokio::test] + async fn chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + // Gemini SSE:每个 data 行是一个完整 generateContent 响应 + let sse = "data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"Hello\"}],\"role\":\"model\"}}]}\n\ + data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\" world\"}],\"role\":\"model\"},\"finishReason\":\"STOP\"}]}\n\n"; + server + .mock("POST", "/models/gemini-2.5-pro:streamGenerateContent") + .match_query(mockito::Matcher::UrlEncoded("alt".into(), "sse".into())) + .match_header("x-goog-api-key", "test-key") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gemini-2.5-pro", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + // 第 1 块文本 "Hello",第 2 块文本 " world" + 结束块 + assert!(chunks.len() >= 2); + let mut content = String::new(); + let mut finish: Option = None; + for chunk in &chunks { + if let Some(c) = chunk + .choices + .first() + .and_then(|d| d.delta.content.as_deref()) + { + content.push_str(c); + } + if let Some(f) = chunk + .choices + .first() + .and_then(|d| d.finish_reason.as_deref()) + { + finish = Some(f.to_string()); + } + } + assert_eq!(content, "Hello world"); + assert_eq!(finish.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_stream_handles_heartbeat_and_empty_lines() { + let mut server = Server::new_async().await; + let sse = ": heartbeat\n\n\ + data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"hi\"}],\"role\":\"model\"}}]}\n\n"; + server + .mock("POST", "/models/m:streamGenerateContent") + .match_query(mockito::Matcher::UrlEncoded("alt".into(), "sse".into())) + .with_status(200) + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 1); + assert_eq!(chunks[0].choices[0].delta.content.as_deref(), Some("hi")); + } + + #[tokio::test] + async fn chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:streamGenerateContent") + .match_query(mockito::Matcher::UrlEncoded("alt".into(), "sse".into())) + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + match adapter.chat_stream(req).await { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn chat_stream_sends_alt_sse_query() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/m:streamGenerateContent") + .match_query(mockito::Matcher::UrlEncoded("alt".into(), "sse".into())) + .with_status(200) + .with_body("data: {\"candidates\":[{\"content\":{\"parts\":[{\"text\":\"x\"}]}}]}\n\n") + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + // ============ image_generate 正常 + 错误路径 ============ + + #[tokio::test] + async fn image_generate_success_parses_inline_data() { + let mut server = Server::new_async().await; + let body = json!({ + "candidates": [{ + "content": { + "parts": [{ + "inlineData": { + "mimeType": "image/png", + "data": "aGVsbG8=" + } + }], + "role": "model" + }, + "finishReason": "STOP" + }] + }); + let mock = server + .mock( + "POST", + "/models/gemini-2.0-flash-exp-image-generation:generateContent", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "contents": [{ "role": "user", "parts": [{ "text": "a cat" }] }], + "generationConfig": { "responseModalities": ["IMAGE", "TEXT"] } + }))) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("gemini-2.0-flash-exp-image-generation", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + assert!(resp.data[0].revised_prompt.is_some()); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_accepts_snake_case_inline_data() { + // 部分 Gemini 兼容端点返回 snake_case 的 inline_data + let mut server = Server::new_async().await; + let body = json!({ + "candidates": [{ + "content": { + "parts": [{ + "inline_data": { + "mime_type": "image/jpeg", + "data": "AAAA" + } + }] + } + }] + }); + server + .mock("POST", "/models/m:generateContent") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("m", "test").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("AAAA")); + } + + #[tokio::test] + async fn image_generate_no_image_returns_empty_data() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(200) + .with_body( + json!({ + "candidates": [{ + "content": { "parts": [{ "text": "no image generated" }] } + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("m", "test").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 0); + } + + #[tokio::test] + async fn image_generate_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:generateContent") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("m", "test").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn image_generate_passes_extra_params() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/m:generateContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "negativePrompt": "blurry" + }))) + .with_status(200) + .with_body(json!({ + "candidates": [{ + "content": { "parts": [{ "inlineData": { "mimeType": "image/png", "data": "x" } }] } + }] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("m", "test") + .extra("negativePrompt", "blurry") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_unsupported_capability() { + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("gemini", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let mut adapter = GeminiAdapter::with_http(http, config); + adapter.capabilities = CapabilitySet::new(); + let req = ImageRequest::builder("m", "test").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ embed 正常 + 错误路径 ============ + + #[tokio::test] + async fn embed_single_input_success() { + let mut server = Server::new_async().await; + let body = json!({ + "embedding": { "values": [0.1, 0.2, 0.3] } + }); + let mock = server + .mock("POST", "/models/text-embedding-004:embedContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "models/text-embedding-004", + "content": { "parts": [{ "text": "hello" }] } + }))) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-004".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!(resp.data[0].index, 0); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2, 0.3]); + } else { + panic!("应为 Float 向量"); + } + assert!(resp.usage.is_none(), "Gemini 嵌入无 usage 统计"); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_single_input_with_dimensions() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/models/m:embedContent") + .match_body(mockito::Matcher::PartialJson(json!({ + "outputDimensionality": 256 + }))) + .with_status(200) + .with_body(json!({"embedding": {"values": [0.1]}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: Some(256), + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let _ = adapter.embed(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_multiple_input_uses_batch_endpoint() { + let mut server = Server::new_async().await; + let body = json!({ + "embeddings": [ + { "values": [0.1, 0.2] }, + { "values": [0.3, 0.4] } + ] + }); + let mock = server + .mock("POST", "/models/text-embedding-004:batchEmbedContents") + .match_body(mockito::Matcher::PartialJson(json!({ + "requests": [ + { "model": "models/text-embedding-004", "content": { "parts": [{ "text": "a" }] } }, + { "model": "models/text-embedding-004", "content": { "parts": [{ "text": "b" }] } } + ] + }))) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-004".into(), + input: EmbedInput::Multiple(vec!["a".into(), "b".into()]), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 2); + assert_eq!(resp.data[0].index, 0); + assert_eq!(resp.data[1].index, 1); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2]); + } else { + panic!("应为 Float 向量"); + } + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/m:embedContent") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn embed_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/models/embed-x:embedContent") + .with_status(404) + .with_body(json!({"error": {"message": "model not found"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "embed-x".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn embed_unsupported_capability() { + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("gemini", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let mut adapter = GeminiAdapter::with_http(http, config); + adapter.capabilities = CapabilitySet::new(); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ list_models 正常 + 错误路径 ============ + + #[tokio::test] + async fn list_models_success_strips_models_prefix() { + let mut server = Server::new_async().await; + let body = json!({ + "models": [ + { + "name": "models/gemini-1.5-pro", + "displayName": "Gemini 1.5 Pro", + "description": "Gemini 1.5 Pro model", + "inputTokenLimit": 2000000 + }, + { + "name": "models/text-embedding-004", + "displayName": "Text Embedding 004" + }, + { + "name": "models/imagen-3.0", + "displayName": "Imagen 3" + } + ] + }); + server + .mock("GET", "/models") + .match_header("x-goog-api-key", "test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + assert_eq!(models[0].id, "gemini-1.5-pro"); + assert_eq!(models[0].name, "Gemini 1.5 Pro"); + assert_eq!(models[0].provider, "gemini"); + assert_eq!(models[0].model_type, ModelType::Chat); + assert_eq!(models[0].max_tokens, Some(2000000)); + assert_eq!(models[1].id, "text-embedding-004"); + // text-embedding-004 不含图像/视频/音频关键字,infer 为 Chat + assert_eq!(models[2].id, "imagen-3.0"); + assert_eq!(models[2].model_type, ModelType::Image); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let mut server = Server::new_async().await; + let body = json!({ + "models": [ + { "name": "models/gemini-1.5-pro", "displayName": "Gemini 1.5 Pro" }, + { "name": "models/imagen-3.0", "displayName": "Imagen 3" } + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "imagen-3.0"); + } + + #[tokio::test] + async fn list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn map_api_error_401_with_message() { + let body = json!({"error": {"code": 401, "message": "API key invalid"}}).to_string(); + let err = GeminiAdapter::map_api_error(401, &body); + match err { + AibridgeError::Authentication { message } => { + assert_eq!(message, "API key invalid"); + } + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn map_api_error_429_is_rate_limit_no_retry_after() { + let body = json!({"error": {"message": "slow down"}}).to_string(); + let err = GeminiAdapter::map_api_error(429, &body); + match err { + AibridgeError::RateLimit { retry_after, .. } => { + assert_eq!(retry_after, None); + } + _ => panic!("应为 RateLimit"), + } + } + + #[test] + fn map_api_error_404_uses_message_as_model() { + let body = json!({"error": {"message": "models/x not found"}}).to_string(); + let err = GeminiAdapter::map_api_error(404, &body); + match err { + AibridgeError::ModelNotFound { model } => { + assert!(model.contains("models/x not found")); + } + _ => panic!("应为 ModelNotFound"), + } + } + + #[test] + fn map_api_error_400_is_validation() { + let body = + json!({"error": {"code": 400, "message": "bad arg", "status": "INVALID_ARGUMENT"}}) + .to_string(); + let err = GeminiAdapter::map_api_error(400, &body); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_api_error_500_is_api() { + let err = GeminiAdapter::map_api_error(503, "service unavailable"); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 503), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_falls_back_to_http_status() { + let err = GeminiAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => { + assert!(message.contains("502")); + } + _ => panic!("应为 Api"), + } + } + + // ============ 参数映射 ============ + + #[test] + fn gemini_mapping_renames_keys() { + let pm = gemini_mapping(); + let mut params = HashMap::new(); + params.insert("max_tokens".to_string(), json!(1000)); + params.insert("top_p".to_string(), json!(0.9)); + params.insert("top_k".to_string(), json!(40)); + params.insert("stop".to_string(), json!(["END"])); + params.insert("temperature".to_string(), json!(0.7)); + let result = pm.apply(¶ms); + assert_eq!( + result.get("maxOutputTokens").and_then(|v| v.as_i64()), + Some(1000) + ); + assert_eq!(result.get("topP").and_then(|v| v.as_f64()), Some(0.9)); + assert_eq!(result.get("topK").and_then(|v| v.as_i64()), Some(40)); + assert!(result.get("stopSequences").is_some()); + // temperature 不在 rename_map,原名透传 + assert_eq!( + result.get("temperature").and_then(|v| v.as_f64()), + Some(0.7) + ); + } + + // ============ 辅助方法测试 ============ + + #[test] + fn convert_image_url_data_uri_to_inline_data() { + let part = GeminiAdapter::convert_image_url("data:image/png;base64,aGVsbG8=").unwrap(); + assert!(part.get("inline_data").is_some()); + assert_eq!(part["inline_data"]["mime_type"].as_str(), Some("image/png")); + assert_eq!(part["inline_data"]["data"].as_str(), Some("aGVsbG8=")); + } + + #[test] + fn convert_image_url_data_uri_jpeg() { + let part = GeminiAdapter::convert_image_url("data:image/jpeg;base64,AAAA").unwrap(); + assert_eq!( + part["inline_data"]["mime_type"].as_str(), + Some("image/jpeg") + ); + } + + #[test] + fn convert_image_url_remote_url_to_file_data() { + let part = GeminiAdapter::convert_image_url("https://example.com/img.png").unwrap(); + assert!(part.get("file_data").is_some()); + assert_eq!( + part["file_data"]["file_uri"].as_str(), + Some("https://example.com/img.png") + ); + } + + #[test] + fn parse_gemini_error_message_extracts_error_message() { + let body = json!({"error": {"code": 400, "message": "bad arg"}}).to_string(); + let msg = parse_gemini_error_message(&body, 400); + assert_eq!(msg, "bad arg"); + } + + #[test] + fn parse_gemini_error_message_falls_back_to_http_status() { + let msg = parse_gemini_error_message("", 500); + assert_eq!(msg, "HTTP 500"); + } + + #[test] + fn map_finish_reason_stop_variants() { + assert_eq!(map_finish_reason("STOP"), "stop"); + assert_eq!(map_finish_reason("MAX_TOKENS"), "length"); + assert_eq!(map_finish_reason("SAFETY"), "content_filter"); + assert_eq!(map_finish_reason("RECITATION"), "stop"); + assert_eq!(map_finish_reason("UNKNOWN"), "stop"); + } + + #[test] + fn parse_usage_metadata_extracts_tokens() { + let v = json!({ + "promptTokenCount": 10, + "candidatesTokenCount": 5, + "totalTokenCount": 15 + }); + let usage = parse_usage_metadata(&v).unwrap(); + assert_eq!(usage.prompt_tokens, 10); + assert_eq!(usage.completion_tokens, 5); + assert_eq!(usage.total_tokens, 15); + } + + #[test] + fn parse_usage_metadata_falls_back_to_sum() { + let v = json!({ "promptTokenCount": 10, "candidatesTokenCount": 5 }); + let usage = parse_usage_metadata(&v).unwrap(); + assert_eq!(usage.total_tokens, 15); + } + + #[test] + fn parse_usage_metadata_none_without_prompt() { + let v = json!({ "candidatesTokenCount": 5 }); + assert!(parse_usage_metadata(&v).is_none()); + } + + #[tokio::test] + async fn base_url_uses_config_when_provided() { + let config = ProviderConfig::from_options( + "gemini", + ClientOptions::builder() + .api_key("k") + .base_url("https://custom.example.com/v1beta") + .build(), + ); + let adapter = GeminiAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.example.com/v1beta"); + } + + #[tokio::test] + async fn base_url_falls_back_to_default_when_missing() { + let config = + ProviderConfig::from_options("gemini", ClientOptions::builder().api_key("k").build()); + let adapter = GeminiAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_GEMINI_BASE_URL); + } + + #[tokio::test] + async fn requires_api_key_is_true() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn capabilities_contains_chat_image_embed() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::Embedding)); + } + + #[tokio::test] + async fn provider_metadata_correct() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + assert_eq!(adapter.provider_type(), "gemini"); + assert_eq!(adapter.provider_name(), "Google Gemini"); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = GeminiAdapter::new(ProviderConfig::from_options( + "gemini", + ClientOptions::builder().api_key("k").build(), + )) + .unwrap(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn unsupported_video_returns_unsupported_capability() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let req = crate::model::video::VideoRequest::builder("m", "prompt").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn unsupported_speech_returns_unsupported_capability() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let req = crate::model::audio::SpeechRequest::builder("m", "hi", "v1").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn unsupported_transcribe_returns_unsupported_capability() { + let server = Server::new_async().await; + let adapter = make_adapter(&server); + let req = + crate::model::audio::TranscribeRequest::builder("m", FileInput::path("/tmp/a.mp3")) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } +} From 0200d4118a3f42f8deb9cda07d263c5ad9c69c2c Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 19:57:15 +0800 Subject: [PATCH 17/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1.2=20agnes=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/agnes.rs | 1477 ++++++++++++++++++++ 1 file changed, 1477 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/agnes.rs diff --git a/crates/aibridge-core/src/adapters/agnes.rs b/crates/aibridge-core/src/adapters/agnes.rs new file mode 100644 index 0000000..bb280a0 --- /dev/null +++ b/crates/aibridge-core/src/adapters/agnes.rs @@ -0,0 +1,1477 @@ +//! Agnes AI 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/agnes.py`。 +//! +//! Agnes AI 是 OpenAI 兼容协议:chat / chat_stream / image_generate / embed / +//! list_models 全部复用 `OpenAiCompatAdapter` 地基实现;视频生成(Video V2.0) +//! 是 Agnes 特有协议(POST /videos 创建任务 + GET /videos/{task_id} 轮询), +//! 本模块独立实现 `video_create` / `video_poll`。 +//! +//! 设计要点(与设计文档第 10 节一致): +//! - 组合(而非继承)`OpenAiCompatAdapter`:AgnesAdapter 内部持有一个 compat 实例, +//! 实现 `Adapter` trait 时把 chat/image/embed/list_models 委托给 compat +//! - 视频协议独立实现:参考 Python `agnes.py` 的 `video_create` / `video_poll` +//! - `list_models` 实时拉取(v1.1.0 特性):直接复用 compat 的实现,无需 override +//! - `requires_api_key = true`:Agnes 需要 API Key 认证 +//! +//! 能力声明(对齐 Python `agnes.py` 的 `supported_capabilities`): +//! CHAT / CHAT_STREAM / IMAGE_GENERATE / VIDEO(VIDEO_TEXT2VIDEO + VIDEO_IMAGE2VIDEO) +//! / EMBEDDING,外加 VISION / TOOL_CALL / JSON_MODE / REASONING 等 chat 子能力。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ChatCompletion, ChatRequest}; +use crate::model::common::{ModelInfo, ModelType, TaskStatus, VoiceInfo}; +use crate::model::image::{ImageRequest, ImageResult}; +use crate::model::options::{EmbedRequest, EmbeddingResult}; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +/// Agnes AI 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`。 +pub const DEFAULT_AGNES_BASE_URL: &str = "https://api.agnes.ai/v1"; + +/// Agnes AI 适配器 +/// +/// 组合 `OpenAiCompatAdapter` 复用 OpenAI 兼容能力,独立实现 Agnes 视频协议。 +/// +/// 构造时传入 `ProviderConfig`,内部据此创建 `OpenAiCompatAdapter` 实例(用于 +/// chat / image / embed / list_models 委托)与独立的 `HttpClient`(用于 Agnes +/// 特有的视频端点)。视频端点不委托 compat,是因为 Agnes 视频协议(POST /videos +/// + GET /videos/{task_id})与 OpenAI 兼容协议差异较大,独立实现更清晰。 +pub struct AgnesAdapter { + /// OpenAI 兼容地基(chat / image / embed / list_models 委托给它) + compat: OpenAiCompatAdapter, + /// 视频端点专用 HTTP 客户端(独立于 compat,避免暴露 compat 私有字段) + http: HttpClient, + /// 视频轮询 URL(Agnes 特有的 /agnesapi 轮询通道,可选) + /// + /// 配置时设置则 `video_poll` 优先走此通道(带 query 参数),失败再回退到 + /// /videos/{task_id}。对应 Python v1 `poll_url` 双路径策略。 + poll_url: Option, + /// Provider 配置(保留引用以便视频方法取 api_key / base_url) + config: ProviderConfig, +} + +impl AgnesAdapter { + /// 创建 Agnes 适配器 + /// + /// - `config`:Provider 配置,`base_url` 为 None 时用 `DEFAULT_AGNES_BASE_URL` 兜底 + /// - 内部构造 `OpenAiCompatAdapter`(委托用)与 `HttpClient`(视频端点用) + pub fn new(config: ProviderConfig) -> Result { + let caps = agnes_capabilities(); + let poll_url = config.poll_url.clone(); + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_AGNES_BASE_URL.to_string()); + + // OpenAiCompatAdapter 内部会从 config.base_url 兜底到传入的默认值 + let compat = OpenAiCompatAdapter::new( + config.clone(), + "agnes", + "Agnes AI", + DEFAULT_AGNES_BASE_URL, + caps, + )?; + + // 视频端点专用 HttpClient:base_url 与 compat 一致,保证 URL 拼接正确 + let http = HttpClient::new( + &ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(), + )?; + + Ok(Self { + compat, + http, + poll_url, + config, + }) + } + + /// 用显式的 HttpClient 构造(测试用,可注入 mockito 后端) + /// + /// 接受一个已构造的 `OpenAiCompatAdapter` 与 `HttpClient`,便于测试时 + /// 复用 mockito 注入逻辑。两者应指向同一 base_url。 + #[cfg(test)] + pub fn with_compat( + compat: OpenAiCompatAdapter, + http: HttpClient, + config: ProviderConfig, + poll_url: Option, + ) -> Self { + Self { + compat, + http, + poll_url, + config, + } + } + + /// Agnes 支持的能力集合 + /// + /// 对齐 Python v1 `agnes.py` 的 `supported_capabilities`。 + fn capabilities_set(&self) -> &CapabilitySet { + self.compat.capabilities_set() + } + + /// API key(可能为空,但 Agnes 要求非空) + fn api_key(&self) -> Option<&str> { + self.config.api_key.as_deref() + } + + /// base_url(已合并 config 与默认值) + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 发送带认证的 POST JSON 请求,并用 OpenAI 错误映射处理响应 + /// + /// 视频端点同样走 OpenAI 兼容的错误体格式 `{"error": {"message": "..."}}`, + /// 复用 `OpenAiCompatAdapter::map_api_error` 统一映射。 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带认证的 GET 请求(带可选 query 参数),并用 OpenAI 错误映射处理响应 + async fn get_authed_json(&self, path: &str, query: &[(&str, &str)]) -> Result { + let url = self.url(path); + let mut req = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key().unwrap_or("")); + for (k, v) in query { + req = req.query(&[(*k, *v)]); + } + let resp = req.send().await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + // ==================== 视频协议(Agnes 特有) ==================== + + /// 构造视频创建请求体 + /// + /// 对应 Python v1 `agnes.py:video_create` 的 body 构造逻辑: + /// - 基础字段:model / prompt + /// - 可选参数:width / height / num_frames / frame_rate / mode / seed / negative_prompt + /// - 参考图像:按 mode 分流到 extra_body + /// - keyframes 模式 + ≥2 图:extra_body.keyframes = { start, end } + /// - multiimage 模式 + ≥2 图:extra_body.image = [urls] + /// - 其余 ≥1 图:extra_body.image = url + fn build_video_body(req: &VideoRequest) -> Value { + let mut body = json!({ + "model": req.model, + "prompt": req.prompt, + }); + + if req.width != 0 { + body["width"] = json!(req.width); + } + if req.height != 0 { + body["height"] = json!(req.height); + } + if let Some(n) = req.num_frames { + body["num_frames"] = json!(n); + } + if req.frame_rate != 0 { + body["frame_rate"] = json!(req.frame_rate); + } + // mode 序列化为小写字符串(text2video / image2video / keyframes / multiimage) + body["mode"] = json!(serde_json::to_value(req.mode).unwrap_or(json!("text2video"))); + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + if let Some(np) = &req.negative_prompt { + body["negative_prompt"] = json!(np); + } + + // 参考图像分流到 extra_body + if !req.reference_images.is_empty() { + // 提取参考图像的 URL 字符串(FileInput::Url 取其 URL,其余类型序列化为字符串) + let urls: Vec = req + .reference_images + .iter() + .map(file_input_to_string) + .collect(); + let extra_body: Value = match req.mode { + crate::model::common::VideoMode::Keyframes if urls.len() >= 2 => json!({ + "keyframes": { "start": urls[0], "end": urls[urls.len() - 1] } + }), + crate::model::common::VideoMode::Multiimage if urls.len() >= 2 => json!({ + "image": urls + }), + _ => json!({ "image": urls[0] }), + }; + body["extra_body"] = extra_body; + } + + // extra 透传:合并到顶层 + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 解析视频创建响应 → VideoTask + /// + /// 对应 Python v1 `agnes.py:video_create` 的响应解析: + /// - task_id:优先 `id`,回退 `video_id`,再回退生成 ID + /// - status:默认 "pending" + /// - created_at:默认当前时间戳 + fn parse_video_task(value: &Value, model: &str) -> VideoTask { + let task_id = value + .get("id") + .and_then(|v| v.as_str()) + .or_else(|| value.get("video_id").and_then(|v| v.as_str())) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("vid")); + let status = value + .get("status") + .and_then(|v| v.as_str()) + .map(parse_task_status) + .unwrap_or(TaskStatus::Pending); + let created_at = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + VideoTask { + task_id, + model: model.to_string(), + status, + created_at, + } + } + + /// 解析视频轮询响应 → VideoStatus + /// + /// 对应 Python v1 `agnes.py:video_poll` 的响应解析: + /// - status:默认 "pending" + /// - video_url:优先顶层 `video_url`,回退 `output.video_url` + /// - progress / error / created / updated 透传 + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let status = value + .get("status") + .and_then(|v| v.as_str()) + .map(parse_task_status) + .unwrap_or(TaskStatus::Pending); + let video_url = value + .get("video_url") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("video_url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + let progress = value + .get("progress") + .and_then(|v| v.as_u64()) + .map(|p| p as u32); + let error = value + .get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let created_at = value.get("created").and_then(|v| v.as_u64()); + let updated_at = value.get("updated").and_then(|v| v.as_u64()); + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress, + error, + created_at, + updated_at, + } + } +} + +/// 将 FileInput 转为可放入请求体的字符串表示 +/// +/// Agnes 视频协议的参考图像字段期望 URL 字符串。 +/// - `Url(s)` → s +/// - `Base64(s)` → s(直接透传 base64 字符串) +/// - `Path(s)` → s(透传路径,由服务端解析) +/// - `Bytes(_)` → 序列化为 JSON 字符串(极少用,保留兜底) +fn file_input_to_string(f: &crate::model::image::FileInput) -> String { + use crate::model::image::FileInput; + match f { + FileInput::Url(s) | FileInput::Path(s) | FileInput::Base64(s) => s.clone(), + FileInput::Bytes(b) => serde_json::to_string(b).unwrap_or_default(), + } +} + +/// 解析任务状态字符串 → TaskStatus +/// +/// 容忍服务端返回的多种写法(success/succeeded/failed/failure 等)。 +fn parse_task_status(s: &str) -> TaskStatus { + let lower = s.to_lowercase(); + match lower.as_str() { + "success" | "succeeded" | "completed" | "done" => TaskStatus::Success, + "failed" | "failure" | "error" => TaskStatus::Failed, + "processing" | "running" | "generating" => TaskStatus::Processing, + _ => TaskStatus::Pending, + } +} + +/// 构造 Agnes 支持的能力集合 +fn agnes_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + // 对话能力 + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::ToolCall); + caps.insert(Capabilities::JsonMode); + caps.insert(Capabilities::Reasoning); + // 图像能力 + caps.insert(Capabilities::ImageGenerate); + // 视频能力 + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + // 嵌入能力 + caps.insert(Capabilities::Embedding); + caps +} + +#[async_trait] +impl Adapter for AgnesAdapter { + fn provider_type(&self) -> &str { + "agnes" + } + + fn provider_name(&self) -> &str { + "Agnes AI" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在构造时已创建,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 走 Drop 释放,无需显式关闭 + Ok(()) + } + + /// 文本对话(委托给 OpenAiCompatAdapter) + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话(委托给 OpenAiCompatAdapter) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 图像生成(委托给 OpenAiCompatAdapter) + async fn image_generate(&self, req: ImageRequest) -> Result { + self.compat.image_generate(req).await + } + + /// 创建视频生成任务(Agnes 特有协议) + /// + /// POST /videos,body 直传参数(model/prompt/width/height/.../extra_body)。 + /// 参考 Python v1 `agnes.py:video_create`。 + async fn video_create(&self, req: VideoRequest) -> Result { + // 能力校验:必须支持视频生成 + if !self + .capabilities_set() + .contains(&Capabilities::VideoGenerate) + { + return Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: agnes)", Capabilities::VideoGenerate.as_str()), + }); + } + let body = Self::build_video_body(&req); + let value = self.post_authed_json("videos", &body).await?; + Ok(Self::parse_video_task(&value, &req.model)) + } + + /// 查询视频任务状态(Agnes 特有协议) + /// + /// 优先走 poll_url(Agnes 特有的 /agnesapi 轮询通道,带 video_id/model_name query), + /// 网络错误时回退到 GET /videos/{task_id}。对应 Python v1 `agnes.py:video_poll` + /// 的双路径策略。 + async fn video_poll(&self, task_id: &str, model: &str) -> Result { + if !self + .capabilities_set() + .contains(&Capabilities::VideoGenerate) + { + return Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: agnes)", Capabilities::VideoGenerate.as_str()), + }); + } + + // 优先走 poll_url 通道(若配置) + if let Some(poll_url) = &self.poll_url { + let query: [(&str, &str); 2] = [("video_id", task_id), ("model_name", model)]; + match self.get_authed_json(poll_url, &query).await { + Ok(value) => return Ok(Self::parse_video_status(&value, task_id)), + Err(AibridgeError::Network(_) | AibridgeError::Timeout) => { + // 网络错误时回退到 /videos/{task_id}(旧版兼容路径) + let value = self + .get_authed_json(&format!("videos/{task_id}"), &[]) + .await?; + return Ok(Self::parse_video_status(&value, task_id)); + } + Err(e) => return Err(e), // 4xx 等 API 错误直接抛出,不回退 + } + } + + // 无 poll_url:直接走 /videos/{task_id} + let value = self + .get_authed_json(&format!("videos/{task_id}"), &[]) + .await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 文本嵌入(委托给 OpenAiCompatAdapter) + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + /// 模型列表(实时拉取,委托给 OpenAiCompatAdapter) + /// + /// v1.1.0 特性:GET /models 实时拉取,不再使用硬编码列表。 + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + /// 列出可用音色(Agnes 不支持音频能力) + async fn list_voices(&self, _language: Option<&str>) -> Result> { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: agnes)", Capabilities::ListVoices.as_str()), + }) + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::adapters::openai_compat::DEFAULT_OPENAI_BASE_URL; + use crate::config::ClientOptions; + use crate::http::HttpClient; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::common::VideoMode; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + use mockito::Server; + use std::collections::HashMap; + + /// 构造测试用 AgnesAdapter(指向 mockito server) + /// + /// 为 compat(chat/image/embed/list_models)与视频端点各创建一个 HttpClient, + /// 两者指向同一 mockito server,保证 URL 拼接一致。 + fn make_adapter(server: &Server, poll_url: Option) -> AgnesAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("agnes", opts); + let compat_http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let video_http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let caps = agnes_capabilities(); + let compat = + OpenAiCompatAdapter::with_http(compat_http, config.clone(), "agnes", "Agnes AI", caps); + AgnesAdapter::with_compat(compat, video_http, config, poll_url) + } + + /// 不带 poll_url 的便捷构造 + fn make_adapter_no_poll(server: &Server) -> AgnesAdapter { + make_adapter(server, None) + } + + // ============ Adapter trait 基本属性 ============ + + #[tokio::test] + async fn provider_type_is_agnes() { + let server = Server::new_async().await; + let adapter = make_adapter_no_poll(&server); + assert_eq!(adapter.provider_type(), "agnes"); + } + + #[tokio::test] + async fn provider_name_is_agnes_ai() { + let server = Server::new_async().await; + let adapter = make_adapter_no_poll(&server); + assert_eq!(adapter.provider_name(), "Agnes AI"); + } + + #[tokio::test] + async fn requires_api_key_is_true() { + let server = Server::new_async().await; + let adapter = make_adapter_no_poll(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn capabilities_includes_chat_image_video_embed() { + let server = Server::new_async().await; + let adapter = make_adapter_no_poll(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + assert!(caps.contains(&Capabilities::Embedding)); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_adapter_no_poll(&server); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ chat 委托(正常 + 错误路径) ============ + + #[tokio::test] + async fn chat_delegates_to_compat() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "chatcmpl-agnes-1", + "object": "chat.completion", + "created": 1700000000, + "model": "agnes-chat", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello from Agnes!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 4, "total_tokens": 9} + }); + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = ChatRequest::builder("agnes-chat", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "chatcmpl-agnes-1"); + assert_eq!(resp.model, "agnes-chat"); + assert_eq!(resp.choices.len(), 1); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("Hello from Agnes!") + ); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 9); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = ChatRequest::builder("agnes-chat", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body( + json!({"error": {"message": "Rate limit exceeded", "retry_after": 2.0}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = ChatRequest::builder("agnes-chat", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::RateLimit { retry_after, .. } => { + assert_eq!(retry_after, Some(2.0)); + } + _ => panic!("应为 RateLimit"), + } + } + + // ============ image 委托 ============ + + #[tokio::test] + async fn image_generate_delegates_to_compat() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{ + "url": "https://example.com/agnes-img.png", + "revised_prompt": "a cute cat" + }] + }); + server + .mock("POST", "/images/generations") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = ImageRequest::builder("agnes-image", "a cat") + .size("1024x1024") + .build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://example.com/agnes-img.png") + ); + } + + #[tokio::test] + async fn image_generate_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter_no_poll(&server); + let req = ImageRequest::builder("agnes-image", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ video_create 正常 + 错误路径 ============ + + #[tokio::test] + async fn video_create_success_returns_task() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "vid-task-1", + "status": "pending", + "created": 1700000000 + }); + let mock = server + .mock("POST", "/videos") + .match_header("authorization", "Bearer test-key") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "seedance-2.0", + "prompt": "a cat walking", + "mode": "text2video" + }))) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "a cat walking").build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + assert_eq!(task.task_id, "vid-task-1"); + assert_eq!(task.model, "seedance-2.0"); + assert_eq!(task.status, TaskStatus::Pending); + assert_eq!(task.created_at, 1700000000); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_falls_back_to_video_id_field() { + let mut server = Server::new_async().await; + // 服务端返回 video_id 而非 id + let body = json!({"video_id": "vid-alt-1", "status": "processing"}); + server + .mock("POST", "/videos") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "prompt").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "vid-alt-1"); + assert_eq!(task.status, TaskStatus::Processing); + } + + #[tokio::test] + async fn video_create_sends_width_height_and_mode() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "seedance-2.0", + "prompt": "a cat", + "width": 1920, + "height": 1080, + "mode": "image2video", + "seed": 42 + }))) + .with_status(200) + .with_body(json!({"id": "t-1", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "a cat") + .width(1920) + .height(1080) + .mode(VideoMode::Image2Video) + .seed(42) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_image2video_single_image_in_extra_body() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos") + .match_body(mockito::Matcher::PartialJson(json!({ + "mode": "image2video", + "extra_body": {"image": "https://example.com/ref.png"} + }))) + .with_status(200) + .with_body(json!({"id": "t-1"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/ref.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_keyframes_mode_uses_start_end() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos") + .match_body(mockito::Matcher::PartialJson(json!({ + "mode": "keyframes", + "extra_body": { + "keyframes": { + "start": "https://example.com/start.png", + "end": "https://example.com/end.png" + } + } + }))) + .with_status(200) + .with_body(json!({"id": "t-1"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Keyframes) + .reference_images(vec![ + FileInput::url("https://example.com/start.png"), + FileInput::url("https://example.com/end.png"), + ]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_multiimage_mode_uses_image_array() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos") + .match_body(mockito::Matcher::PartialJson(json!({ + "mode": "multiimage", + "extra_body": { + "image": ["https://example.com/a.png", "https://example.com/b.png"] + } + }))) + .with_status(200) + .with_body(json!({"id": "t-1"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Multiimage) + .reference_images(vec![ + FileInput::url("https://example.com/a.png"), + FileInput::url("https://example.com/b.png"), + ]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_create_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "prompt").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn video_create_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "prompt").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn video_create_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = VideoRequest::builder("seedance-2.0", "prompt").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ video_poll 正常 + 错误路径 ============ + + #[tokio::test] + async fn video_poll_success_via_videos_endpoint() { + let mut server = Server::new_async().await; + let body = json!({ + "status": "success", + "video_url": "https://example.com/video.mp4", + "progress": 100, + "created": 1700000000, + "updated": 1700000100 + }); + let mock = server + .mock("GET", "/videos/vid-task-1") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let status = adapter + .video_poll("vid-task-1", "seedance-2.0") + .await + .expect("video_poll 应成功"); + assert_eq!(status.task_id, "vid-task-1"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/video.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert_eq!(status.created_at, Some(1700000000)); + assert_eq!(status.updated_at, Some(1700000100)); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_poll_reads_video_url_from_output_field() { + let mut server = Server::new_async().await; + // video_url 嵌套在 output 对象中 + let body = json!({ + "status": "success", + "output": {"video_url": "https://example.com/nested.mp4"} + }); + server + .mock("GET", "/videos/t-2") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let status = adapter.video_poll("t-2", "seedance-2.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/nested.mp4") + ); + } + + #[tokio::test] + async fn video_poll_processing_status() { + let mut server = Server::new_async().await; + let body = json!({"status": "processing", "progress": 45}); + server + .mock("GET", "/videos/t-3") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let status = adapter.video_poll("t-3", "seedance-2.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(45)); + assert!(status.video_url.is_none()); + } + + #[tokio::test] + async fn video_poll_failed_status_with_error() { + let mut server = Server::new_async().await; + let body = json!({"status": "failed", "error": "content policy violation"}); + server + .mock("GET", "/videos/t-4") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let status = adapter.video_poll("t-4", "seedance-2.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("content policy violation")); + } + + #[tokio::test] + async fn video_poll_prefers_poll_url_when_configured() { + let mut server = Server::new_async().await; + // poll_url 走 /agnesapi/video/query 路径,带 query 参数 + let body = json!({"status": "success", "video_url": "https://example.com/v.mp4"}); + let mock = server + .mock("GET", "/agnesapi/video/query") + .match_query(mockito::Matcher::AllOf(vec![ + mockito::Matcher::UrlEncoded("video_id".into(), "vid-x".into()), + mockito::Matcher::UrlEncoded("model_name".into(), "seedance-2.0".into()), + ])) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server, Some("agnesapi/video/query".to_string())); + let status = adapter.video_poll("vid-x", "seedance-2.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/v.mp4") + ); + mock.assert_async().await; + } + + #[tokio::test] + async fn video_poll_does_not_fallback_on_api_error() { + let mut server = Server::new_async().await; + // poll_url 返回 500(API 错误,非 Network/Timeout),不应回退,直接抛 Api 错误 + server + .mock("GET", "/agnesapi/video/query") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + // /videos/{id} 不应被调用:mockito 用 expect(0) 断言 0 次命中 + server + .mock("GET", "/videos/vid-y") + .expect(0) + .create_async() + .await; + + let adapter = make_adapter(&server, Some("agnesapi/video/query".to_string())); + // 500 是 API 错误(非 Network/Timeout),直接抛出,不回退 + let err = adapter + .video_poll("vid-y", "seedance-2.0") + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::Api { .. })); + } + + #[tokio::test] + async fn video_poll_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/t-5") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let err = adapter.video_poll("t-5", "seedance-2.0").await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn video_poll_error_404_returns_api() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/t-missing") + .with_status(404) + .with_body(json!({"error": {"message": "task not found"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let err = adapter + .video_poll("t-missing", "seedance-2.0") + .await + .unwrap_err(); + // 404 在 OpenAI 错误映射里走 ModelNotFound + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + // ============ embed 委托 ============ + + #[tokio::test] + async fn embed_delegates_to_compat() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [{"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}], + "model": "agnes-embed", + "usage": {"prompt_tokens": 2, "total_tokens": 2} + }); + server + .mock("POST", "/embeddings") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let req = EmbedRequest { + model: "agnes-embed".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!(resp.model, "agnes-embed"); + } + + #[tokio::test] + async fn embed_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter_no_poll(&server); + let req = EmbedRequest { + model: "agnes-embed".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ list_models 委托(实时拉取) ============ + + #[tokio::test] + async fn list_models_pulls_from_models_endpoint() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"id": "agnes-chat", "object": "model", "created": 1700000000, "owned_by": "agnes"}, + {"id": "seedream-4.0", "object": "model", "created": 1700000000, "owned_by": "agnes"}, + {"id": "seedance-2.0", "object": "model", "created": 1700000000, "owned_by": "agnes"} + ] + }); + let mock = server + .mock("GET", "/models") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + // 类型推断 + assert_eq!(models[0].id, "agnes-chat"); + assert_eq!(models[0].model_type, ModelType::Chat); + assert_eq!(models[1].model_type, ModelType::Image); + assert_eq!(models[2].model_type, ModelType::Video); + // provider 字段填充为 agnes + assert_eq!(models[0].provider, "agnes"); + mock.assert_async().await; + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "agnes-chat", "object": "model", "created": 1, "owned_by": "agnes"}, + {"id": "seedance-2.0", "object": "model", "created": 1, "owned_by": "agnes"} + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter_no_poll(&server); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert_eq!(videos.len(), 1); + assert_eq!(videos[0].id, "seedance-2.0"); + } + + #[tokio::test] + async fn list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter_no_poll(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ list_voices 不支持 ============ + + #[tokio::test] + async fn list_voices_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_adapter_no_poll(&server); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ 构造与 base_url 兜底 ============ + + #[tokio::test] + async fn new_uses_default_base_url_when_missing() { + let config = + ProviderConfig::from_options("agnes", ClientOptions::builder().api_key("k").build()); + let adapter = AgnesAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_AGNES_BASE_URL); + } + + #[tokio::test] + async fn new_uses_config_base_url_when_provided() { + let config = ProviderConfig::from_options( + "agnes", + ClientOptions::builder() + .api_key("k") + .base_url("https://custom.agnes.example.com/v1") + .build(), + ); + let adapter = AgnesAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.agnes.example.com/v1"); + } + + #[tokio::test] + async fn new_preserves_poll_url() { + let config = ProviderConfig::from_options( + "agnes", + ClientOptions::builder() + .api_key("k") + .poll_url("agnesapi/video/query") + .build(), + ); + let adapter = AgnesAdapter::new(config).unwrap(); + assert_eq!(adapter.poll_url.as_deref(), Some("agnesapi/video/query")); + } + + // ============ 内部解析函数单元测试 ============ + + #[test] + fn parse_task_status_recognizes_variants() { + assert_eq!(parse_task_status("success"), TaskStatus::Success); + assert_eq!(parse_task_status("SUCCEEDED"), TaskStatus::Success); + assert_eq!(parse_task_status("completed"), TaskStatus::Success); + assert_eq!(parse_task_status("failed"), TaskStatus::Failed); + assert_eq!(parse_task_status("FAILURE"), TaskStatus::Failed); + assert_eq!(parse_task_status("error"), TaskStatus::Failed); + assert_eq!(parse_task_status("processing"), TaskStatus::Processing); + assert_eq!(parse_task_status("running"), TaskStatus::Processing); + assert_eq!(parse_task_status("queued"), TaskStatus::Pending); + assert_eq!(parse_task_status("unknown"), TaskStatus::Pending); + } + + #[test] + fn parse_video_task_uses_id_first() { + let value = json!({"id": "t-id", "video_id": "t-vid", "status": "pending"}); + let task = AgnesAdapter::parse_video_task(&value, "m"); + assert_eq!(task.task_id, "t-id"); + assert_eq!(task.model, "m"); + assert_eq!(task.status, TaskStatus::Pending); + } + + #[test] + fn parse_video_task_falls_back_to_video_id() { + let value = json!({"video_id": "t-vid", "status": "success"}); + let task = AgnesAdapter::parse_video_task(&value, "m"); + assert_eq!(task.task_id, "t-vid"); + assert_eq!(task.status, TaskStatus::Success); + } + + #[test] + fn parse_video_task_generates_id_when_missing() { + let value = json!({"status": "pending"}); + let task = AgnesAdapter::parse_video_task(&value, "m"); + assert!(task.task_id.starts_with("vid_"), "应生成 vid_ 前缀 ID"); + } + + #[test] + fn parse_video_status_extracts_top_level_video_url() { + let value = json!({"status": "success", "video_url": "https://x.com/v.mp4"}); + let s = AgnesAdapter::parse_video_status(&value, "t-1"); + assert_eq!(s.video_url.as_deref(), Some("https://x.com/v.mp4")); + } + + #[test] + fn parse_video_status_extracts_nested_output_video_url() { + let value = json!({ + "status": "success", + "output": {"video_url": "https://x.com/nested.mp4"} + }); + let s = AgnesAdapter::parse_video_status(&value, "t-1"); + assert_eq!(s.video_url.as_deref(), Some("https://x.com/nested.mp4")); + } + + #[test] + fn build_video_body_text2video_minimal() { + let req = VideoRequest::builder("seedance-2.0", "a cat").build(); + let body = AgnesAdapter::build_video_body(&req); + assert_eq!(body["model"], "seedance-2.0"); + assert_eq!(body["prompt"], "a cat"); + assert_eq!(body["mode"], "text2video"); + // 无参考图像时不应有 extra_body + assert!(body.get("extra_body").is_none()); + } + + #[test] + fn build_video_body_includes_all_optional_params() { + let req = VideoRequest::builder("seedance-2.0", "a cat") + .width(1920) + .height(1080) + .num_frames(120) + .frame_rate(30) + .seed(42) + .negative_prompt("blurry") + .build(); + let body = AgnesAdapter::build_video_body(&req); + assert_eq!(body["width"], 1920); + assert_eq!(body["height"], 1080); + assert_eq!(body["num_frames"], 120); + assert_eq!(body["frame_rate"], 30); + assert_eq!(body["seed"], 42); + assert_eq!(body["negative_prompt"], "blurry"); + } + + #[test] + fn build_video_body_extra_passthrough() { + let req = VideoRequest::builder("seedance-2.0", "a cat") + .extra("custom_param", "custom_value") + .build(); + let body = AgnesAdapter::build_video_body(&req); + assert_eq!(body["custom_param"], "custom_value"); + } + + #[test] + fn build_video_body_keyframes_with_two_images() { + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Keyframes) + .reference_images(vec![ + FileInput::url("https://x.com/start.png"), + FileInput::url("https://x.com/end.png"), + ]) + .build(); + let body = AgnesAdapter::build_video_body(&req); + assert_eq!( + body["extra_body"]["keyframes"]["start"], + "https://x.com/start.png" + ); + assert_eq!( + body["extra_body"]["keyframes"]["end"], + "https://x.com/end.png" + ); + } + + #[test] + fn build_video_body_multiimage_with_two_images() { + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Multiimage) + .reference_images(vec![ + FileInput::url("https://x.com/a.png"), + FileInput::url("https://x.com/b.png"), + ]) + .build(); + let body = AgnesAdapter::build_video_body(&req); + let images = body["extra_body"]["image"].as_array().unwrap(); + assert_eq!(images.len(), 2); + assert_eq!(images[0], "https://x.com/a.png"); + } + + #[test] + fn build_video_body_image2video_single_image() { + let req = VideoRequest::builder("seedance-2.0", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://x.com/ref.png")]) + .build(); + let body = AgnesAdapter::build_video_body(&req); + assert_eq!(body["extra_body"]["image"], "https://x.com/ref.png"); + } + + #[test] + fn file_input_to_string_handles_variants() { + assert_eq!( + file_input_to_string(&FileInput::url("https://x.com/a.png")), + "https://x.com/a.png" + ); + assert_eq!( + file_input_to_string(&FileInput::path("/tmp/a.png")), + "/tmp/a.png" + ); + assert_eq!( + file_input_to_string(&FileInput::base64("aGVsbG8=")), + "aGVsbG8=" + ); + } + + #[test] + fn agnes_capabilities_contains_expected_set() { + let caps = agnes_capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + assert!(caps.contains(&Capabilities::Embedding)); + // Agnes 不声明音频能力 + assert!(!caps.contains(&Capabilities::AudioSpeech)); + assert!(!caps.contains(&Capabilities::AudioTranscribe)); + } + + /// 编译期断言:DEFAULT_AGNES_BASE_URL 与 openai_compat 默认值不同 + #[test] + fn agnes_default_base_url_differs_from_openai() { + assert_ne!(DEFAULT_AGNES_BASE_URL, DEFAULT_OPENAI_BASE_URL); + assert_eq!(DEFAULT_AGNES_BASE_URL, "https://api.agnes.ai/v1"); + } +} From bf5cb1db7106f6f3b55c5ae9cb56a622b89a87f7 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:02:15 +0800 Subject: [PATCH 18/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?1=20=E6=94=B6=E5=B0=BE=20=E6=B3=A8=E5=86=8C=E5=9B=9B=20MVP=20?= =?UTF-8?q?=E9=80=82=E9=85=8D=E5=99=A8=E5=88=B0=E5=B7=A5=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 62 +++++++++------------ crates/aibridge-core/src/adapters/mod.rs | 12 ++++ crates/aibridge-core/src/client.rs | 11 ++-- 3 files changed, 41 insertions(+), 44 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index 2c35910..4c7e976 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -8,7 +8,11 @@ //! - 阶段 0.4 暂只占位分支(返 ProviderNotFound),具体适配器阶段 1 起填充 use crate::adapter::Adapter; +use crate::adapters::agnes::AgnesAdapter; use crate::adapters::echo::EchoAdapter; +use crate::adapters::gemini::GeminiAdapter; +use crate::adapters::openai::OpenAiAdapter; +use crate::adapters::volcengine_cv::VolcengineCvAdapter; use crate::config::ProviderConfig; use crate::error::{AibridgeError, Result}; @@ -44,17 +48,15 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ /// `echo` 为阶段 0.6 管线验证用 mock 适配器,已实现; /// 其余 provider 阶段 0.4 占位(返 ProviderNotFound),阶段 1 起填充。 pub fn create_adapter(config: ProviderConfig) -> Result> { - let provider = config.provider_type.as_str(); - match provider { + let provider = config.provider_type.clone(); + match provider.as_str() { // Echo(Mock)适配器:阶段 0.6 管线验证用,已实现 "echo" => Ok(Box::new(EchoAdapter::new())), - // 阶段 1 MVP 适配器(阶段 1.0 起填充实际构造逻辑) - "openai" | "agnes" | "volcengine_cv" | "gemini" => { - // TODO(阶段 1): 引入 adapters::openai::OpenAiAdapter 等具体实现 - Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 1 待实现)"), - }) - } + // 阶段 1 MVP 适配器:阶段 1.0 已实现具体构造逻辑 + "openai" => Ok(Box::new(OpenAiAdapter::new(config)?)), + "agnes" => Ok(Box::new(AgnesAdapter::new(config)?)), + "volcengine_cv" => Ok(Box::new(VolcengineCvAdapter::new(config)?)), + "gemini" => Ok(Box::new(GeminiAdapter::new(config)?)), // 阶段 2 适配器占位 "azure" | "anthropic" @@ -103,43 +105,29 @@ mod tests { } #[test] - fn create_openai_returns_pending_for_phase0() { - let result = create_adapter(config_for("openai")); - // 阶段 0.4:占位返 ProviderNotFound(待阶段 1 实现) - assert!(matches!( - result, - Err(AibridgeError::ProviderNotFound { .. }) - )); - if let Err(AibridgeError::ProviderNotFound { provider }) = result { - assert!(provider.contains("阶段 1")); - } + fn create_openai_returns_adapter() { + // 阶段 1:工厂已能构造真实 OpenAiAdapter(仅校验构造成功,不触发 HTTP) + let adapter = create_adapter(config_for("openai")).expect("工厂应能创建 openai 适配器"); + assert_eq!(adapter.provider_type(), "openai"); } #[test] - fn create_agnes_returns_pending() { - let result = create_adapter(config_for("agnes")); - assert!(matches!( - result, - Err(AibridgeError::ProviderNotFound { .. }) - )); + fn create_agnes_returns_adapter() { + let adapter = create_adapter(config_for("agnes")).expect("工厂应能创建 agnes 适配器"); + assert_eq!(adapter.provider_type(), "agnes"); } #[test] - fn create_volcengine_cv_returns_pending() { - let result = create_adapter(config_for("volcengine_cv")); - assert!(matches!( - result, - Err(AibridgeError::ProviderNotFound { .. }) - )); + fn create_volcengine_cv_returns_adapter() { + let adapter = + create_adapter(config_for("volcengine_cv")).expect("工厂应能创建 volcengine_cv 适配器"); + assert_eq!(adapter.provider_type(), "volcengine_cv"); } #[test] - fn create_gemini_returns_pending() { - let result = create_adapter(config_for("gemini")); - assert!(matches!( - result, - Err(AibridgeError::ProviderNotFound { .. }) - )); + fn create_gemini_returns_adapter() { + let adapter = create_adapter(config_for("gemini")).expect("工厂应能创建 gemini 适配器"); + assert_eq!(adapter.provider_type(), "gemini"); } #[test] diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 3c7d8cd..5128209 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -15,3 +15,15 @@ pub mod echo; /// OpenAI 兼容协议适配器地基:阶段 1.0 实现,为 openai/agnes 等子适配器提供共享基础 pub mod openai_compat; + +/// Agnes 适配器:阶段 1 MVP,OpenAI 兼容协议的子适配器 +pub mod agnes; + +/// Gemini 适配器:阶段 1 MVP,Google Gemini 独立协议 +pub mod gemini; + +/// OpenAI 适配器:阶段 1 MVP,OpenAI 官方协议 +pub mod openai; + +/// 火山引擎 CV 适配器:阶段 1 MVP,火山引擎视觉/视频生成协议 +pub mod volcengine_cv; diff --git a/crates/aibridge-core/src/client.rs b/crates/aibridge-core/src/client.rs index 7daaeee..717320f 100644 --- a/crates/aibridge-core/src/client.rs +++ b/crates/aibridge-core/src/client.rs @@ -168,13 +168,10 @@ mod tests { } #[test] - fn new_with_key_reaches_factory_and_returns_provider_not_found() { - // 阶段 0.4:工厂占位返 ProviderNotFound - let result = Client::new("openai", ClientOptions::builder().api_key("sk-xxx").build()); - assert!(matches!( - result, - Err(AibridgeError::ProviderNotFound { .. }) - )); + fn new_with_key_reaches_factory_and_creates_adapter() { + // 阶段 1:工厂已能构造真实 OpenAiAdapter(仅校验 Client 构建成功,不触发 HTTP) + let _client = Client::new("openai", ClientOptions::builder().api_key("sk-xxx").build()) + .expect("应成功创建 openai Client(api_key 已提供,工厂应返 Ok)"); } #[test] From 68954d223a9db6322d18e1f979996ff9104cac51 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:10:55 +0800 Subject: [PATCH 19/55] =?UTF-8?q?feat:=20=E9=98=B6=E6=AE=B51.5=20=E8=B7=A8?= =?UTF-8?q?=E8=AF=AD=E8=A8=80=E4=B8=80=E8=87=B4=E6=80=A7=E6=B5=8B=E8=AF=95?= =?UTF-8?q?=20+=20=E9=94=99=E8=AF=AF=20code=20=E7=BB=9F=E4=B8=80?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 五语言错误 code 对齐 Rust aibridge-core error.rs 的 AibridgeError::code() 实际返回值(authentication_error / rate_limit_error / timeout_error 等): - Go (bindings/go/error.go):errCodeAuthentication/RateLimit/Network/Timeout 四个 常量补齐 _error 后缀,更新注释示例 - .NET (bindings/dotnet/AIBridge/AibridgeException.cs):MapByCode switch 与 Authentication/RateLimit/Timeout 三个子类构造的 code 字符串统一为 Rust 实际值 - Python/Node/JVM 原本已对齐,无需修改 跨语言一致性验证(echo adapter): - 四语言 hello world(Python/Node/Go/JVM)chat/stream/speech 输出语义一致 (chat="hello [echo]",stream 3 chunk 拼接="hello [echo]",speech=15 字节) - 四语言未知 provider 错误探针 code 均为 provider_not_found - .NET 因本机无 dotnet SDK 仅做代码级对齐确认 - 一致性测试文档与探针脚本置于 tests/consistency/ --- bindings/dotnet/AIBridge/AibridgeException.cs | 23 ++- bindings/go/error.go | 14 +- tests/consistency/ErrorProbe.java | 41 +++++ .../consistency/cross_language_consistency.md | 155 ++++++++++++++++++ tests/consistency/error_probe_go.go | 48 ++++++ tests/consistency/error_probe_node.js | 32 ++++ tests/consistency/error_probe_python.py | 40 +++++ 7 files changed, 337 insertions(+), 16 deletions(-) create mode 100644 tests/consistency/ErrorProbe.java create mode 100644 tests/consistency/cross_language_consistency.md create mode 100644 tests/consistency/error_probe_go.go create mode 100644 tests/consistency/error_probe_node.js create mode 100644 tests/consistency/error_probe_python.py diff --git a/bindings/dotnet/AIBridge/AibridgeException.cs b/bindings/dotnet/AIBridge/AibridgeException.cs index e92d70f..278c363 100644 --- a/bindings/dotnet/AIBridge/AibridgeException.cs +++ b/bindings/dotnet/AIBridge/AibridgeException.cs @@ -6,17 +6,22 @@ namespace AIBridge; // 异常体系 // // 对应设计文档第 9 节 .NET 异常映射:AibridgeException + 子类。 -// 子类与 core AibridgeError 枚举变体一一对应(rate_limit/authentication/...)。 +// 子类与 core AibridgeError 枚举变体一一对应(rate_limit_error/authentication_error/...)。 // // FFI 错误模型(设计文档 7.4):aibridge_status_t 返回码 + aibridge_last_error() // 线程局部 JSON:{"code":"...","message":"...","details":...,"retryable":bool}。 // Client 在 FFI 失败后同线程读取 last_error 转存,按 code 映射为子类异常。 +// +// 注意:code 字符串必须与 aibridge-core error.rs 的 AibridgeError::code() 实际返回值 +// 完全一致(authentication_error / rate_limit_error / validation_error / model_not_found / +// api_error / network_error / timeout_error / unsupported_capability / provider_not_found / +// voice_not_available / service_unavailable),以保证五语言跨绑定错误码统一。 // ============================================================================ /// AIBridge 异常基类。 public class AibridgeException : Exception { - /// 错误码字符串(如 "rate_limit"、"authentication")。 + /// 错误码字符串(如 "rate_limit_error"、"authentication_error")。 public string Code { get; } /// 是否可重试。 @@ -84,18 +89,18 @@ internal static AibridgeException FromStatus(int status, string? lastErrorJson) }; } - /// 按 last_error 的 code 字段映射子类(core AibridgeError 变体名)。 + /// 按 last_error 的 code 字段映射子类(与 core AibridgeError::code() 对齐)。 private static AibridgeException? MapByCode(string code, string message, bool retryable, JsonElement? details) { return code switch { - "authentication" => new AuthenticationException(message, retryable, details), - "rate_limit" => new RateLimitException(message, retryable, details), + "authentication_error" => new AuthenticationException(message, retryable, details), + "rate_limit_error" => new RateLimitException(message, retryable, details), "validation_error" => new ValidationException(message, retryable, details), "model_not_found" => new ModelNotFoundException(message, retryable, details), "api_error" => new ApiException(message, retryable, details), "network_error" => new NetworkException(message, retryable, details), - "timeout" => new TimeoutException_(message, retryable, details), + "timeout_error" => new TimeoutException_(message, retryable, details), "unsupported_capability" => new UnsupportedCapabilityException(message, retryable, details), "provider_not_found" => new ProviderNotFoundException(message, retryable, details), "voice_not_available" => new VoiceNotAvailableException(message, retryable, details), @@ -110,7 +115,7 @@ internal static AibridgeException FromStatus(int status, string? lastErrorJson) public class AuthenticationException : AibridgeException { public AuthenticationException(string msg, bool retryable = false, JsonElement? details = null) - : base(msg, "authentication", retryable, details) { } + : base(msg, "authentication_error", retryable, details) { } } public class RateLimitException : AibridgeException @@ -119,7 +124,7 @@ public class RateLimitException : AibridgeException public double? RetryAfter { get; } public RateLimitException(string msg, bool retryable = true, JsonElement? details = null) - : base(msg, "rate_limit", retryable, details) + : base(msg, "rate_limit_error", retryable, details) { // 尝试从 details.retry_after 读取 if (details.HasValue && details.Value.TryGetProperty("retry_after", out JsonElement r) @@ -158,7 +163,7 @@ public NetworkException(string msg, bool retryable = true, JsonElement? details public class TimeoutException_ : AibridgeException { public TimeoutException_(string msg, bool retryable = true, JsonElement? details = null) - : base(msg, "timeout", retryable, details) { } + : base(msg, "timeout_error", retryable, details) { } } public class UnsupportedCapabilityException : AibridgeException diff --git a/bindings/go/error.go b/bindings/go/error.go index 636d56f..703cbad 100644 --- a/bindings/go/error.go +++ b/bindings/go/error.go @@ -7,7 +7,7 @@ // // var err error = ... // if ae, ok := err.(AibridgeError); ok { -// fmt.Println(ae.Code()) // "rate_limit" 等 +// fmt.Println(ae.Code()) // "rate_limit_error" 等 // } package aibridge @@ -21,7 +21,7 @@ import ( // 对应设计文档 9.3 节:Go 用 error + 类型断言接口 type AibridgeError interface{ Code() string } type AibridgeError interface { error - Code() string // 错误码(如 "rate_limit"、"validation_error") + Code() string // 错误码(如 "rate_limit_error"、"validation_error") Retryable() bool // 是否可重试 Message() string // 原始错误消息 } @@ -84,15 +84,15 @@ func newFfiError(msg string) AibridgeError { } } -// 常见错误码常量(与 aibridge.h 的 AIBRIDGE_ERR_* 宏对齐,供需要按码判断时使用) +// 常见错误码常量(与 aibridge-core error.rs 的 AibridgeError::code() 实际返回值对齐) const ( - errCodeAuthentication = "authentication" - errCodeRateLimit = "rate_limit" + errCodeAuthentication = "authentication_error" + errCodeRateLimit = "rate_limit_error" errCodeValidation = "validation_error" errCodeModelNotFound = "model_not_found" errCodeAPI = "api_error" - errCodeNetwork = "network" - errCodeTimeout = "timeout" + errCodeNetwork = "network_error" + errCodeTimeout = "timeout_error" errCodeUnsupportedCapability = "unsupported_capability" errCodeProviderNotFound = "provider_not_found" errCodeVoiceNotAvailable = "voice_not_available" diff --git a/tests/consistency/ErrorProbe.java b/tests/consistency/ErrorProbe.java new file mode 100644 index 0000000..d853737 --- /dev/null +++ b/tests/consistency/ErrorProbe.java @@ -0,0 +1,41 @@ +import io.aibridge.AibridgeException; +import io.aibridge.Client; + +/** + * 跨语言错误一致性探针(JVM):未知 provider 必须返回 code == "provider_not_found"。 + * + *

JVM 绑定从 FFI last_error JSON 解析 code 字段({@link AibridgeException#getCode()}), + * code 直接来自 core {@code AibridgeError::code()}。 + * + *

编译运行(classpath 含 jvm classes + jna jar): + *

{@code
+ * javac -cp : tests/consistency/ErrorProbe.java
+ * java -Djna.library.path=target/debug -cp ::tests/consistency ErrorProbe
+ * }
+ * + * 退出码 0 表示通过,1 表示失败。 + */ +public class ErrorProbe { + + private static final String EXPECTED = "provider_not_found"; + + public static void main(String[] args) { + // 未知 provider + 假 key(跳过 key 校验,触发 ProviderNotFound) + // ClientOptions JSON 字段为 snake_case(core serde 默认) + String configJson = "{\"api_key\":\"dummy-key\"}"; + try { + new Client("nonexistent", configJson); + } catch (AibridgeException e) { + if (EXPECTED.equals(e.getCode())) { + System.out.println("[jvm] OK:e.getCode()=\"" + EXPECTED + "\""); + System.out.println("[jvm] message=" + e.getMessage()); + System.exit(0); + } + System.out.println("[jvm] FAIL:期望 code=\"" + EXPECTED + "\",实际 code=\"" + e.getCode() + "\""); + System.out.println("[jvm] message=" + e.getMessage()); + System.exit(1); + } + System.out.println("[jvm] FAIL:未抛出任何异常"); + System.exit(1); + } +} diff --git a/tests/consistency/cross_language_consistency.md b/tests/consistency/cross_language_consistency.md new file mode 100644 index 0000000..6477ac4 --- /dev/null +++ b/tests/consistency/cross_language_consistency.md @@ -0,0 +1,155 @@ +# AIBridge 阶段1.5 跨语言一致性测试 + +> 分支:`feat/aibridge-rust-rewrite` +> 日期:2026-07-07 +> 范围:五语言错误 code 统一 + 四语言(Python/Node/Go/JVM)hello world 一致性验证。 + +## 1. 目标 + +1. **错误 code 统一**:五个语言绑定(Python / Node / Go / JVM / .NET)的错误码字符串完全对齐 Rust `aibridge-core/src/error.rs` 中 `AibridgeError::code()` 的实际返回值。 +2. **跨语言一致性**:用 echo adapter(`provider="echo"`,免认证)跑 Python/Node/Go/JVM 的 hello world,断言四种能力的输出语义一致: + - `chat`:`choices[0].message.content` == `"hello [echo]"` + - `chat_stream`:3 个 chunk,拼接内容 == `"hello [echo]"` + - `speech`:`audio_data` 长度 == 15 + - 错误:未知 provider 返回 code == `"provider_not_found"` +3. **.NET**:dotnet 工具链未安装,仅做代码级错误 code 对齐校验,不跑 hello。 + +## 2. 基准:Rust `AibridgeError::code()` + +来源:`crates/aibridge-core/src/error.rs`(`code()` 方法,第 255-269 行)。 + +| 变体 | code() 返回值 | +|------|--------------| +| Authentication | `authentication_error` | +| RateLimit | `rate_limit_error` | +| Validation | `validation_error` | +| ModelNotFound | `model_not_found` | +| Api | `api_error` | +| Network | `network_error` | +| Timeout | `timeout_error` | +| UnsupportedCapability | `unsupported_capability` | +| ProviderNotFound | `provider_not_found` | +| VoiceNotAvailable | `voice_not_available` | +| ServiceUnavailable | `service_unavailable` | + +FFI 层(`aibridge-ffi`)把上述 code 写入线程局部 `last_error` JSON 的 `code` 字段;运行时各绑定均从该字段读取,故运行时 code 天然对齐 Rust。本阶段修复的是各绑定**源码内硬编码的 code 字符串常量 / switch case / 子类构造传入的 code**,使其与 Rust 一致(避免兜底分支或子类 `Code` 属性与实际 last_error code 不一致)。 + +## 3. 错误 code 统一前后对比 + +### 3.1 Python(`crates/aibridge-python/src/lib.rs`) + +- **机制**:`map_error()` 用 `err.code()` 拼接消息 `"[{code}] {err}"`,异常类层级映射。 +- **改动**:无(已对齐)。 + +### 3.2 Node(`crates/aibridge-node/src/lib.rs` + `lib.js`) + +- **机制**:`map_error()` 用 `err.code()` 编码 `[code] message`;`lib.js` 的 `withCode()` 解析出 `err.code` 属性。 +- **改动**:无(已对齐)。 + +### 3.3 JVM(`bindings/jvm/src/main/java/io/aibridge/AibridgeException.java`) + +- **机制**:`CODE_*` 常量 + `mapToException()` switch。常量已对齐 Rust。 +- **改动**:无(已对齐)。 + +### 3.4 Go(`bindings/go/error.go`)— **已修改** + +| 常量 | 修改前 | 修改后 | +|------|--------|--------| +| `errCodeAuthentication` | `"authentication"` | `"authentication_error"` | +| `errCodeRateLimit` | `"rate_limit"` | `"rate_limit_error"` | +| `errCodeNetwork` | `"network"` | `"network_error"` | +| `errCodeTimeout` | `"timeout"` | `"timeout_error"` | +| 其余常量 | (已对齐) | (不变) | + +同时更新文件头注释与接口注释里的示例(`"rate_limit"` → `"rate_limit_error"`)。 +注:这些常量当前未被业务代码引用(运行时 code 来自 FFI last_error JSON),但作为文档/兜底用途须与 Rust 一致。 + +### 3.5 .NET(`bindings/dotnet/AIBridge/AibridgeException.cs`)— **已修改** + +| 位置 | 修改前 | 修改后 | +|------|--------|--------| +| `MapByCode` switch:authentication | `"authentication"` | `"authentication_error"` | +| `MapByCode` switch:rate_limit | `"rate_limit"` | `"rate_limit_error"` | +| `MapByCode` switch:timeout | `"timeout"` | `"timeout_error"` | +| `AuthenticationException` 构造 code | `"authentication"` | `"authentication_error"` | +| `RateLimitException` 构造 code | `"rate_limit"` | `"rate_limit_error"` | +| `TimeoutException_` 构造 code | `"timeout"` | `"timeout_error"` | + +修改前 `MapByCode` 的 switch case 与子类构造传入的 code 字符串不一致(前者部分用短名,后者部分用短名),导致 `MapByCode` 命中后构造的子类 `Code` 属性与 last_error 的 code 不一致(例如 last_error code=`"rate_limit_error"`,但 `MapByCode` 旧 case `"rate_limit"` 不命中 → 落到状态码兜底分支 → `RateLimitException.Code` 仍为旧 `"rate_limit"`)。修改后两者完全对齐 Rust。 + +同时更新类头注释与 `Code` 属性文档注释。 + +## 4. 四语言 hello world 一致性验证 + +环境:Python 3.14 + maturin 1.14 / Node 25.5 / Go 1.26 / OpenJDK 21。 +echo adapter:`chat` 回显最后一条 user 消息 + `" [echo]"`;`chat_stream` 产 3 chunk(role / 前半段 `"hello "` / 后半段 `"[echo]"` + `finish_reason="stop"`);`speech` 返 15 字节 `b"mock-audio-data"`。 + +| 语言 | 脚本 | chat content | stream chunk 数 | stream 拼接 | speech 字节数 | 结果 | +|------|------|-------------|----------------|-------------|--------------|------| +| Python | `examples/hello_python.py` | `hello [echo]` | 3 | `hello [echo]` | 15 | ✅ 通过 | +| Node | `examples/hello_node.js` | `hello [echo]` | 3 | `hello [echo]` | 15 | ✅ 通过 | +| Go | `bindings/go/example/hello.go` | `hello [echo]` | 3 | `hello [echo]` | 15 | ✅ 通过 | +| JVM | `bindings/jvm` (`./gradlew run`,`Hello.java`) | `hello [echo]` | 3 | `hello [echo]` | 15 | ✅ 通过 | + +四语言在 chat / chat_stream / speech 三项能力的输出语义完全一致。 + +## 5. 未知 provider 错误一致性验证 + +探针脚本(`tests/consistency/error_probe_*`):构造 `Client(provider="nonexistent", api_key="dummy-key")`(假 key 跳过 key 校验,触发 `ProviderNotFound`),断言错误 code == `"provider_not_found"`。 + +| 语言 | 探针脚本 | 错误对象 | code 取值方式 | 实际 code | 结果 | +|------|---------|---------|--------------|----------|------| +| Python | `error_probe_python.py` | `ProviderNotFoundError` | 消息前缀 `[code]` | `provider_not_found` | ✅ | +| Node | `error_probe_node.js` | `Error` | `err.code` 属性 | `provider_not_found` | ✅ | +| Go | `error_probe_go.go` | `aibridge.AibridgeError` | `ae.Code()` | `provider_not_found` | ✅ | +| JVM | `ErrorProbe.java` | `AibridgeException` | `e.getCode()` | `provider_not_found` | ✅ | + +四语言错误 code 完全一致,均等于 Rust `AibridgeError::ProviderNotFound.code()` 返回的 `"provider_not_found"`。 + +### 5.1 运行方式 + +```bash +# 前置:cargo build -p aibridge-ffi(产出 target/debug/libaibridge.dylib,须为 ffi 版本而非 PyO3 版本) + +# Python(需先 maturin develop -m crates/aibridge-python/Cargo.toml) +.venv/bin/python tests/consistency/error_probe_python.py + +# Node(需先在 crates/aibridge-node 下 npm install && npm run build) +node tests/consistency/error_probe_node.js + +# Go +cd bindings/go +DYLD_LIBRARY_PATH=../../target/debug CGO_ENABLED=1 go run ../../tests/consistency/error_probe_go.go + +# JVM(需先 ./gradlew classes 编译) +JVM_CLASSES=bindings/jvm/build/classes/java/main +JNA_JAR=$(find ~/.gradle/caches -name 'jna-5.15.0.jar' | head -1) +JACKSON_DATABIND=$(find ~/.gradle/caches -name 'jackson-databind-2.18.1.jar' | head -1) +JACKSON_CORE=$(find ~/.gradle/caches -name 'jackson-core-2.18.1.jar' | head -1) +JACKSON_ANN=$(find ~/.gradle/caches -name 'jackson-annotations-2.18.1.jar' | head -1) +mkdir -p /tmp/aibridge_probe +javac -cp "$JVM_CLASSES:$JNA_JAR" -d /tmp/aibridge_probe tests/consistency/ErrorProbe.java +DYLD_LIBRARY_PATH=target/debug java -Djna.library.path=target/debug \ + -cp "$JVM_CLASSES:$JNA_JAR:$JACKSON_DATABIND:$JACKSON_CORE:$JACKSON_ANN:/tmp/aibridge_probe" ErrorProbe +``` + +## 6. .NET 代码级对齐确认(未跑 hello) + +dotnet 工具链未安装,仅做代码审查。`bindings/dotnet/AIBridge/AibridgeException.cs` 修改后: + +- `MapByCode` switch 的 11 个 case 字符串与 Rust `code()` 11 个返回值逐一对应。 +- 11 个子类构造函数传入的 code 字符串与各自 `MapByCode` case 一致。 +- 状态码兜底分支(`FromStatus`)在 last_error 缺失时按 `AibridgeStatus` 映射子类,子类 `Code` 属性现已对齐 Rust。 + +## 7. 结论 + +- 五语言错误 code 已统一对齐 Rust `AibridgeError::code()` 实际返回值(Go / .NET 修改,Python / Node / JVM 原本对齐)。 +- 四语言(Python / Node / Go / JVM)hello world 在 chat / chat_stream / speech / 未知 provider 错误四项上输出语义完全一致。 +- .NET 代码级错误 code 已对齐(dotnet 工具链缺失,未跑运行时验证)。 + +## 8. 遗留问题 + +1. **dylib 产物名冲突**:`aibridge-ffi` 与 `aibridge-python` 的 `[lib] name = "aibridge"`,两者 `cargo build` 产物都落到 `target/debug/libaibridge.dylib`,互相覆盖。`maturin develop` 后再跑 Go/JVM 会因 dylib 变成 PyO3 版本(无 C FFI 符号)而链接失败。**当前规避**:跑 Go/JVM 前先 `touch crates/aibridge-ffi/src/lib.rs && cargo build -p aibridge-ffi` 重建 ffi 版本。**建议后续**:给 aibridge-python 的 cdylib 改名(如 `aibridge_python`)或在 workspace 层隔离产物路径。 +2. **.NET 运行时验证缺失**:本机未装 dotnet SDK,.NET 仅做代码审查。建议在装了 dotnet 的环境补跑 `dotnet run` hello 与错误探针。 +3. **Go 常量未被引用**:`bindings/go/error.go` 的 `errCode*` 常量当前无业务代码引用(运行时 code 来自 FFI last_error JSON)。建议后续在 Go 侧加一个 `MapByCode` 风格的子类型断言(或保留常量作文档)。 +4. **一致性测试未纳入 CI**:本阶段的探针脚本为手动运行。建议后续接入 CI matrix(五语言并行跑 hello + 错误探针)。 diff --git a/tests/consistency/error_probe_go.go b/tests/consistency/error_probe_go.go new file mode 100644 index 0000000..371dd90 --- /dev/null +++ b/tests/consistency/error_probe_go.go @@ -0,0 +1,48 @@ +// 跨语言错误一致性探针(Go):未知 provider 必须返回 code == "provider_not_found"。 +// +// Go 绑定从 FFI last_error JSON 解析出 AibridgeError.Code()(parseErrorJSON), +// code 字段直接来自 core AibridgeError::code()。 +// +// 运行: +// +// cd bindings/go +// DYLD_LIBRARY_PATH=../../target/debug CGO_ENABLED=1 go run ../../tests/consistency/error_probe_go.go +// +// 退出码 0 表示通过,1 表示失败。 +package main + +import ( + "fmt" + "os" + + aibridge "github.com/aibridge/aibridge-go" +) + +const expected = "provider_not_found" + +func main() { + // 未知 provider + 假 key(跳过 key 校验,触发 ProviderNotFound) + // ClientOptions JSON 字段为 snake_case(core serde 默认) + opts := `{"api_key":"dummy-key"}` + _, err := aibridge.NewClient("nonexistent", &opts) + if err == nil { + fmt.Println("[go] FAIL:未抛出任何错误") + os.Exit(1) + } + + // 类型断言取 Code() + ae, ok := err.(aibridge.AibridgeError) + if !ok { + fmt.Printf("[go] FAIL:错误不是 AibridgeError 接口,实际 %T: %v\n", err, err) + os.Exit(1) + } + + if ae.Code() == expected { + fmt.Printf("[go] OK:err.Code()=%q\n", expected) + fmt.Printf("[go] message=%q\n", ae.Message()) + os.Exit(0) + } + fmt.Printf("[go] FAIL:期望 code=%q,实际 code=%q\n", expected, ae.Code()) + fmt.Printf("[go] message=%q\n", ae.Message()) + os.Exit(1) +} diff --git a/tests/consistency/error_probe_node.js b/tests/consistency/error_probe_node.js new file mode 100644 index 0000000..a1a3802 --- /dev/null +++ b/tests/consistency/error_probe_node.js @@ -0,0 +1,32 @@ +'use strict'; + +// 跨语言错误一致性探针(Node):未知 provider 必须返回 err.code === "provider_not_found"。 +// +// lib.js 的 withCode 会从 Rust map_error 编码的 "[code] message" 中解析出 .code 属性。 +// 退出码 0 表示通过,1 表示失败。 + +const { Client } = require('../../crates/aibridge-node'); + +const EXPECTED = 'provider_not_found'; + +function main() { + try { + // 未知 provider + 假 key(跳过 key 校验,触发 ProviderNotFound) + // ClientOptions 字段为 snake_case(core serde 默认,无 rename_all) + // eslint-disable-next-line no-new + new Client('nonexistent', { api_key: 'dummy-key' }); + } catch (err) { + if (err.code === EXPECTED) { + console.log(`[node] OK:err.code=${JSON.stringify(EXPECTED)}`); + console.log(`[node] message=${JSON.stringify(err.message)}`); + return 0; + } + console.log(`[node] FAIL:期望 code=${JSON.stringify(EXPECTED)},实际 code=${JSON.stringify(err.code)}`); + console.log(`[node] message=${JSON.stringify(err.message)}`); + return 1; + } + console.log('[node] FAIL:未抛出任何异常'); + return 1; +} + +process.exit(main()); diff --git a/tests/consistency/error_probe_python.py b/tests/consistency/error_probe_python.py new file mode 100644 index 0000000..3092675 --- /dev/null +++ b/tests/consistency/error_probe_python.py @@ -0,0 +1,40 @@ +"""跨语言错误一致性探针:未知 provider 必须返回 code == "provider_not_found"。 + +构造方式:Client(provider="nonexistent", api_key="dummy-key")。 +- core 侧:api_key 校验通过 → 工厂 create_adapter 返 ProviderNotFound。 +- 各语言绑定把 core AibridgeError::code() 透出为错误对象的 code 属性/字段。 + +期望:四种语言的错误 code 字符串完全一致,均等于 "provider_not_found"。 + +退出码 0 表示通过,1 表示失败。仅打印 code 与结论,不打印栈。 +""" + +import aibridge +from aibridge import Client, ProviderNotFoundError + + +def main() -> int: + expected = "provider_not_found" + try: + # 未知 provider + 假 key(跳过 key 校验,触发 ProviderNotFound) + Client(provider="nonexistent", api_key="dummy-key") + except ProviderNotFoundError as e: + # 异常类型正确,进一步校验消息里的 code 前缀 + msg = str(e) + # core 错误消息格式:[provider_not_found] Provider 不存在: nonexistent + if expected not in msg: + print(f"[python] FAIL:消息未含 code {expected!r},实际 {msg!r}") + return 1 + print(f"[python] OK:异常类型=ProviderNotFoundError,消息含 code={expected!r}") + print(f"[python] message={msg!r}") + return 0 + except Exception as e: # noqa: BLE001 + print(f"[python] FAIL:期望 ProviderNotFoundError,实际 {type(e).__name__}: {e}") + return 1 + print("[python] FAIL:未抛出任何异常") + return 1 + + +if __name__ == "__main__": + import sys + sys.exit(main()) From e5df4f339b6daaff9bd7364e13730bc7777e8530 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:17:02 +0800 Subject: [PATCH 20/55] =?UTF-8?q?fix(aibridge-python):=20dylib=20=E4=BA=A7?= =?UTF-8?q?=E7=89=A9=E5=90=8D=E6=94=B9=E4=B8=BA=20=5Faibridge=20=E9=81=BF?= =?UTF-8?q?=E5=85=8D=E4=B8=8E=20ffi=20=E5=86=B2=E7=AA=81?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit aibridge-ffi 和 aibridge-python 的 [lib] name 都为 "aibridge",导致 cargo build 把两个 cdylib 都输出到 target/debug/libaibridge.dylib 互相覆盖。 maturin develop 后 dylib 变成 PyO3 版本(无 C FFI 符号),Go/JVM 链接失败。 修复(仅改 aibridge-python,ffi 的 [lib] name="aibridge" 保持不变以匹配 Go/JVM/.NET 的 -laibridge 链接): - Cargo.toml: [lib] name 从 "aibridge" 改为 "_aibridge",cdylib 产物变为 lib_aibridge.dylib,与 libaibridge.dylib 分离 - src/lib.rs: #[pymodule] 函数名从 aibridge 改为 _aibridge(生成 PyInit__aibridge 符号与 lib name 一致),加 #[pyo3(name = "aibridge")] 保持模块 Python 名 - pyproject.toml: [tool.maturin] 加 module-name = "aibridge",把 Rust pymodule _aibridge 映射到 Python 顶层模块 aibridge,import aibridge 不受影响 验证: - maturin develop 无 warning,import aibridge 正常 - examples/hello_python.py chat/stream/speech 全部通过 - cargo build -p aibridge-ffi 产 libaibridge.dylib(6 个 C FFI 符号,0 PyInit) - Go hello world 链接 libaibridge 通过 - JVM hello world 通过 - cargo build --workspace 0 warning --- crates/aibridge-python/Cargo.toml | 9 ++++++++- crates/aibridge-python/pyproject.toml | 6 ++++++ crates/aibridge-python/src/lib.rs | 8 +++++++- 3 files changed, 21 insertions(+), 2 deletions(-) diff --git a/crates/aibridge-python/Cargo.toml b/crates/aibridge-python/Cargo.toml index dff497f..73be0f0 100644 --- a/crates/aibridge-python/Cargo.toml +++ b/crates/aibridge-python/Cargo.toml @@ -10,7 +10,14 @@ description = "AIBridge Python 绑定(PyO3,直连 aibridge-core,原生 asy [lib] crate-type = ["cdylib"] -name = "aibridge" +# lib name 用 `_aibridge`(而非 `aibridge`)以避免与 aibridge-ffi 的 +# `libaibridge.dylib` 产物冲突:两者 `[lib] name` 都为 `aibridge` 时,`cargo build` +# 会把 cdylib 都输出到 target/debug/libaibridge.dylib 互相覆盖,导致 `maturin develop` +# 后 dylib 变成 PyO3 版本(无 C FFI 符号),Go/JVM 链接失败。 +# +# Python import 名由 `#[pymodule] fn aibridge` 决定(见 src/lib.rs),与 lib name 无关, +# 故改 lib name 不影响 `import aibridge`。下划线前缀是 Python C 扩展惯例。 +name = "_aibridge" [dependencies] aibridge-core.workspace = true diff --git a/crates/aibridge-python/pyproject.toml b/crates/aibridge-python/pyproject.toml index fe6d650..2bab6da 100644 --- a/crates/aibridge-python/pyproject.toml +++ b/crates/aibridge-python/pyproject.toml @@ -15,4 +15,10 @@ classifiers = [ dynamic = ["version"] [tool.maturin] +# lib name 为 `_aibridge`(见 Cargo.toml [lib] name),故 Rust pymodule 函数名也是 +# `_aibridge`(生成 `PyInit__aibridge` 符号)。通过 `module-name` 把它映射到 Python +# 顶层模块 `aibridge`,使 `import aibridge` 正常工作,同时 cdylib 产物名为 +# `lib_aibridge.dylib`(macOS)/ `lib_aibridge.so`(Linux),不与 aibridge-ffi 的 +# `libaibridge.dylib` 冲突。 +module-name = "aibridge" features = ["pyo3/extension-module"] diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs index e9df520..e453488 100644 --- a/crates/aibridge-python/src/lib.rs +++ b/crates/aibridge-python/src/lib.rs @@ -848,8 +848,14 @@ impl Client { // =========================================================================== /// Python 模块入口:`import aibridge` +/// +/// 函数名为 `_aibridge`(与 `[lib] name = "_aibridge"` 一致),生成 `PyInit__aibridge` +/// 符号。`#[pyo3(name = "aibridge")]` 把模块的 Python 名重写为 `aibridge`,配合 +/// pyproject.toml 的 `module-name = "aibridge"`,使 `import aibridge` 正常工作。 +/// lib name 用下划线前缀以避免与 aibridge-ffi 的 `libaibridge.dylib` 产物冲突。 #[pymodule] -fn aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { +#[pyo3(name = "aibridge")] +fn _aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { let py = m.py(); // 触发全局 runtime 初始化(首次访问 Lazy 即建) From 7b4a79da782ec2a84f27ccd54da6923c450d0dfc Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:23:27 +0800 Subject: [PATCH 21/55] =?UTF-8?q?fix:=20=E9=98=B6=E6=AE=B51=20Python/Node?= =?UTF-8?q?=20=E6=B5=81=E5=BC=8F=E6=A1=A5=E6=8E=A5=E9=87=8D=E6=9E=84?= =?UTF-8?q?=EF=BC=88=E7=9C=9F=E5=AE=9E=20IO=20=E4=B8=8D=E9=98=BB=E5=A1=9E?= =?UTF-8?q?=E4=BA=8B=E4=BB=B6=E5=BE=AA=E7=8E=AF=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit Python(crates/aibridge-python/src/lib.rs): chat_stream 的 __anext__ 原先用 RUNTIME.block_on 预取 chunk,echo adapter 无 IO 瞬时完成能跑通,但真实 adapter(openai/agnes/gemini 的 reqwest IO) 会阻塞 asyncio 事件循环。 重构为 PyO3 0.28 内置 Coroutine 桥接模式: - __anext__ 同步返回 pyo3::coroutine::Coroutine(可 await 的 Python 对象), 其包裹的 Rust future 在全局 tokio runtime 上 spawn 消费 core ChatStream。 - 真实 IO 在 tokio worker 线程执行,asyncio 线程仅 await JoinHandle(Pending 时让出,其他协程可运行),不阻塞事件循环。 - Coroutine 的 AsyncioWaker 自动通过 asyncio.Future + call_soon_threadsafe 把 "chunk 就绪" 通知桥接回 asyncio 事件循环(PyO3 内置实现,无需手写 asyncio loop 引用),await 期间让出线程。 - 移除手写的 NextAwaitable(__await__/__iter__/__next__ 协议)及模块注册, 改用 Coroutine 内置的 send/__next__/__await__。 - GIL:future 在 tokio 上 await stream 期间不持 GIL,仅在拿到 chunk 后 Python::attach 构造 pyclass。 - 流结束 future 返回 Err(StopAsyncIteration);chunk 出错返回对应 AibridgeError 子类;正常 chunk 返回 Ok(chunk) → StopIteration(chunk)。 Node(crates/aibridge-node/src/lib.rs): 经核查无需重构。当前 mpsc channel + #[napi] async fn next() 模式不阻塞 libuv 事件循环:napi 2.x 用多线程 tokio runtime,#[napi] async fn 的 future 与 spawn 的消费 task 都在 tokio worker 线程调度,recv().await Pending 时让出 tokio 线程,libuv 事件循环自由运行。JS await Promise 让 其他任务运行,与 Python 端机制等价。 验证: - cargo build --workspace 通过 - cargo clippy -p aibridge-python -p aibridge-node -- -D warnings 无警告 - maturin develop + python examples/hello_python.py:echo chat_stream 3 chunk 拼接 'hello [echo]' 正确 - napi build + node examples/hello_node.js:echo chatStream 3 chunk 拼接 'hello [echo]' 正确 - 真实 IO(openai/agnes/gemini)未验证(无 API key),但重构基于 "tokio task 消费 stream + asyncio await Future" 的非阻塞推理 --- crates/aibridge-python/src/lib.rs | 148 +++++++++++++++--------------- 1 file changed, 76 insertions(+), 72 deletions(-) diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs index e453488..2ea809f 100644 --- a/crates/aibridge-python/src/lib.rs +++ b/crates/aibridge-python/src/lib.rs @@ -10,9 +10,12 @@ //! tokio 上执行,PyO3 协程通过 await JoinHandle 拿结果。echo adapter 无网络也走 //! 同一路径,保持一致。 //! - `PyClient` 持有 `Arc>`,支持 `start`/`close` 可变操作。 -//! - 流式:`chat_stream` 把 core `ChatStream` 通过 `tokio::sync::Mutex` 封装进 -//! `ChatStreamIterator`,`__anext__` 同步返回 awaitable(PyO3 0.28 slot 不支持 -//! async fn),awaitable 内抛 `StopIteration(chunk)`/`StopAsyncIteration`。 +//! - 流式:`chat_stream` 把 core `ChatStream`(`BoxStream`)通过 `tokio::sync::Mutex` +//! 封装进 `ChatStreamIterator`。`__anext__` 返回 PyO3 内置的 +//! [`pyo3::coroutine::Coroutine`],其包裹的 Rust future 在 tokio runtime 上 +//! `spawn` 消费 stream(真实 IO 在 tokio worker 线程,不阻塞 asyncio 事件循环)。 +//! `Coroutine` 的 waker 通过 `asyncio.Future` + `call_soon_threadsafe` 把 +//! "chunk 就绪"通知回 asyncio 事件循环,await 期间让出线程给其他协程。 // PyO3 0.28 的 `#[pymethods]` 宏展开会用到 Rust 1.77+ 稳定的语法(如 let 链), // 与 workspace MSRV 1.75 冲突。该 lint 针对宏生成代码,非手写代码,故整体允许。 @@ -22,7 +25,8 @@ use std::sync::Arc; use futures::StreamExt; use pyo3::prelude::*; -use pyo3::types::PyBytes; +use pyo3::coroutine::Coroutine; +use pyo3::types::{PyBytes, PyString}; use tokio::sync::Mutex; use aibridge_core::adapter::ChatStream as CoreChatStream; @@ -524,11 +528,19 @@ impl SpeechResult { /// `async for chunk in stream:` 每次取一个 `ChatCompletionChunk`,流结束抛 /// `StopAsyncIteration`。 /// -/// 实现说明(PyO3 0.28 限制): -/// `__anext__` slot 不支持 `async fn`,故采用"同步 `__anext__` 返回 awaitable" -/// 模式。`__anext__` 在全局 tokio runtime 上 `block_on` 取下一个 chunk(echo -/// adapter 纯计算瞬时完成;真实 adapter 阶段将改为 asyncio.Future 桥接避免 -/// 阻塞事件循环),把结果封装进 [`NextAwaitable`] 返回。 +/// 实现说明(不阻塞 asyncio 事件循环): +/// `__anext__` 同步返回一个 PyO3 内置的 [`Coroutine`](可 `await` 的 Python 对象), +/// 其包裹的 Rust future 在全局 tokio runtime 上 `spawn` 消费 core `ChatStream`: +/// - 真实 adapter 的 reqwest IO 在 tokio worker 线程执行,asyncio 线程仅 await +/// `JoinHandle`(Pending 时让出,不阻塞事件循环,其他协程可运行)。 +/// - chunk 就绪后,`Coroutine` 的 `AsyncioWaker` 通过 `asyncio.Future` + +/// `call_soon_threadsafe` 把就绪通知调度回 asyncio 事件循环(PyO3 内置实现, +/// 无需手写 loop 引用),`await` 返回 chunk。 +/// - 流结束:future 返回 `Err(StopAsyncIteration)`;取 chunk 出错:返回对应 +/// `AibridgeError` 子类;正常 chunk:`Ok(chunk)` → `StopIteration(chunk)`。 +/// +/// GIL 处理:future 在 tokio 上 await stream 期间不持 GIL(`spawn` 的 task 在 +/// tokio worker 跑),仅在拿到 chunk 后 `Python::with_gil` 构造 Python 对象。 #[pyclass] struct ChatStreamIterator { /// core 流(None 表示已耗尽) @@ -542,73 +554,66 @@ impl ChatStreamIterator { slf } - /// 取下一个 chunk(同步返回 awaitable) + /// 取下一个 chunk(同步返回 `Coroutine`,可 `await`) + /// + /// 返回的 `Coroutine` `await` 后得到 `ChatCompletionChunk`,或抛 + /// `StopAsyncIteration`(流结束)/ 对应 `AibridgeError` 子类(取 chunk 出错)。 /// - /// 返回一个 [`NextAwaitable`],`await` 后得到 `ChatCompletionChunk` 或抛 - /// `StopAsyncIteration`(流结束)。 - fn __anext__(&self, py: Python<'_>) -> PyResult> { + /// 不阻塞事件循环:实际取 chunk 的 future 在 tokio runtime 上推进, + /// `Coroutine` 的 waker 负责把就绪通知桥接回 asyncio 事件循环。 + fn __anext__(&self, py: Python<'_>) -> PyResult> { let inner = self.inner.clone(); - // 在全局 tokio runtime 上同步取下一个 chunk。 - // echo adapter 无 IO,瞬时完成;block_on 在当前线程仅等待 JoinHandle, - // 实际 stream.next() 在 tokio worker 线程执行,不会死锁。 - // py.detach 释放 GIL 期间阻塞,避免长时间持锁。 - let item: Option> = - py.detach(|| RUNTIME.block_on(async move { - let mut guard = inner.lock().await; - match guard.as_mut() { - None => None, - Some(stream) => stream.next().await, - } - })); - let chunk = match item { - None => None, - Some(Ok(c)) => Some(Ok(Py::new(py, ChatCompletionChunk::from_core(c))?)), - Some(Err(e)) => Some(Err(map_error(e))), + // 构造包裹"取下一个 chunk"逻辑的 future。该 future 在 Coroutine 被 + // poll 时推进(poll 发生在 asyncio 线程,持 GIL),但其内部把 stream + // 消费 spawn 到 tokio runtime,await JoinHandle 期间 Pending 让出线程。 + let fut = async move { + // 在 tokio runtime 上消费 stream。spawn 后 await JoinHandle: + // - stream.next()(含真实 reqwest IO)在 tokio worker 线程执行 + // - asyncio 线程仅 poll JoinHandle,Pending 时注册 waker 让出 + let join_result = RUNTIME + .spawn(async move { + let mut guard = inner.lock().await; + match guard.as_mut() { + None => None, + Some(stream) => stream.next().await, + } + }) + .await; + + // JoinError(task panic/取消)→ RuntimeError + let item: Option> = join_result + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!( + "chat_stream 消费任务失败: {e}" + )) + })?; + + // 在 tokio worker 线程拿到 item,需重新进入 GIL 上下文构造 Python 对象。 + // Coroutine future 被 poll 时所在线程(asyncio 线程)已 attached GIL, + // `Python::attach` 在已 attached 线程上直接复用(返回 R),安全构造 pyclass。 + Python::attach(|py: Python<'_>| -> PyResult> { + match item { + // 流结束 → 抛 StopAsyncIteration(await 时终止 async for) + None => Err(pyo3::exceptions::PyStopAsyncIteration::new_err(())), + // 正常 chunk → 返回 chunk 对象(Coroutine 抛 StopIteration(chunk)) + Some(Ok(c)) => { + let chunk = Py::new(py, ChatCompletionChunk::from_core(c))?; + Ok(chunk.into_any()) + } + // 取 chunk 出错 → 抛对应 AibridgeError 子类 + Some(Err(e)) => Err(map_error(e)), + } + }) }; - Py::new(py, NextAwaitable { chunk }) - } -} - -/// `__anext__` 返回的 awaitable -/// -/// 实现 `__await__`/`__iter__`/`__next__` 协议:`await` 时 `__next__` 抛 -/// `StopIteration(chunk)` 返回 chunk,或 `StopAsyncIteration` 表示流结束。 -/// -/// 结果在构造时预计算(由 `__anext__` 的 block_on 完成)。 -#[pyclass] -struct NextAwaitable { - /// 预计算的下一个 chunk(None 表示流已结束) - chunk: Option>>, -} - -#[pymethods] -impl NextAwaitable { - /// `__await__` 返回 self(awaitable 协议) - fn __await__(slf: Py) -> Py { - slf - } - - /// `__iter__` 返回 self(awaitable 兼容迭代器协议) - fn __iter__(slf: Py) -> Py { - slf - } - - /// `__next__` 抛出结果 - /// - /// - 有 chunk:抛 `StopIteration(chunk)`,`await` 得到 chunk - /// - 流结束:抛 `StopAsyncIteration` - /// - 取 chunk 出错:抛对应 AibridgeError 子类 - fn __next__(&self, py: Python<'_>) -> PyResult> { - match &self.chunk { - None => Err(pyo3::exceptions::PyStopAsyncIteration::new_err(())), - Some(Ok(chunk)) => { - // StopIteration(chunk) → await 得到 chunk - Err(pyo3::exceptions::PyStopIteration::new_err((chunk.clone_ref(py),))) - } - Some(Err(e)) => Err(e.clone_ref(py)), - } + // 用 PyO3 内置 Coroutine 包装 future。Coroutine 实现 __await__/__next__/send, + // 可直接被 `await`。其 waker 自动桥接 tokio 唤醒 → asyncio.Future.set_result + // (通过 call_soon_threadsafe),无需手写 asyncio loop 引用。 + let name = PyString::new(py, "ChatStreamIterator.__anext__"); + let coroutine = + pyo3::impl_::coroutine::new_coroutine(&name, Some("ChatStreamIterator"), None, fut); + Py::new(py, coroutine) } } @@ -893,7 +898,6 @@ fn _aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { // 客户端与流式 m.add_class::()?; m.add_class::()?; - m.add_class::()?; // 模块版本 m.add("__version__", aibridge_core::VERSION)?; From 03d3e057453efa79a8527117fe22bd8a946fa538 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:36:27 +0800 Subject: [PATCH 22/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20azure=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/azure.rs | 1362 +++++++++++++++++ .../src/adapters/openai_compat.rs | 33 +- 2 files changed, 1388 insertions(+), 7 deletions(-) create mode 100644 crates/aibridge-core/src/adapters/azure.rs diff --git a/crates/aibridge-core/src/adapters/azure.rs b/crates/aibridge-core/src/adapters/azure.rs new file mode 100644 index 0000000..22706c3 --- /dev/null +++ b/crates/aibridge-core/src/adapters/azure.rs @@ -0,0 +1,1362 @@ +//! Azure OpenAI 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/azure.py`。 +//! +//! Azure OpenAI API 与 OpenAI API 基本兼容(请求体/响应格式相同),但有关键差异: +//! - **Base URL 结构**:`https://{resource}.cognitiveservices.azure.com/openai/deployments/{deployment}` +//! (Python 老版用 `.openai.azure.com`,新版 Azure 已迁移到 `cognitiveservices.azure.com`, +//! 两者都可通过 `config.base_url` 覆盖) +//! - **认证方式**:用 HTTP header `api-key: {key}`,而非 `Authorization: Bearer {key}` +//! - **路径结构**:`POST /openai/deployments/{deployment}/{action}?api-version={version}` +//! - chat:`/chat/completions` +//! - image:`/images/generations` +//! - embed:`/embeddings` +//! - **API 版本**:必须带 `api-version` query 参数(如 `2024-02-15-preview`) +//! - **模型名**:是 deployment name(部署名),非底层模型名 +//! - **list_models**:调 Azure 部署列表 API `GET /openai/deployments?api-version={version}`, +//! 响应里 `id` 是部署名、`model` 是底层模型名 +//! +//! 复用策略:请求体构造(`build_chat_body` / `build_image_body` / `build_embed_body`) +//! 与响应解析(`parse_chat_completion` / `parse_image_result` / `parse_embedding_result`) +//! 以及错误映射(`map_api_error`)与 OpenAI 兼容协议完全一致,故内部持有一个 +//! [`OpenAiCompatAdapter`] 实例,把这部分逻辑委托给它;HTTP 请求层(URL 拼接、 +//! `api-key` header、`api-version` query)由本适配器自行实现,因地基的 +//! `post_authed_json` 写死了 `bearer_auth` 与无 query 的相对路径。 +//! +//! 能力:Chat / ChatStream / ImageGenerate / Embedding(与 Python 老版核心子集一致, +//! Python 还声明了 AUDIO_TRANSCRIBE/TRANSLATE/SPEECH,但阶段 2a 范围仅核心四能力, +//! audio 走 trait 默认实现返 UnsupportedCapability,待阶段 2c audio_adapters 补齐)。 +//! `requires_api_key = true`。 + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::Value; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ChatCompletion, ChatRequest}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::image::{ImageRequest, ImageResult}; +use crate::model::options::{EmbedRequest, EmbeddingResult}; + +/// Azure OpenAI 默认 API 版本 +/// +/// 对应 Python v1 `DEFAULT_API_VERSION = "2024-02-15-preview"`。 +/// Azure 端点必须带 `api-version` query 参数,未配置时用此默认值。 +pub const DEFAULT_AZURE_API_VERSION: &str = "2024-02-15-preview"; + +/// Azure OpenAI 默认资源主机后缀 +/// +/// 老版 Azure 用 `{resource}.openai.azure.com`,新版迁移到 +/// `{resource}.cognitiveservices.azure.com`。本常量仅用于无 `base_url` 时 +/// 由 `resource_name` + `deployment_id` 拼接默认 base_url,采用 Python 老版 +/// 一致的 `openai.azure.com`(用户可通过 `config.base_url` 覆盖为 cognitiveservices)。 +const DEFAULT_AZURE_HOST_SUFFIX: &str = "openai.azure.com"; + +/// Azure 路径前缀(deployments 段之前的固定路径) +const AZURE_PATH_PREFIX: &str = "/openai/deployments"; + +/// Azure 适配器 +/// +/// 持有 HTTP 客户端、Provider 配置、解析后的 Azure 专用字段(resource_name / +/// deployment_id / api_version)以及一个 [`OpenAiCompatAdapter`] 实例(用于 +/// 复用 OpenAI 兼容的请求体构造与响应解析)。 +/// +/// 构造时即解析 base_url 与 Azure 字段,`start` / `close` 为空操作 +/// (HTTP 客户端在 `new` 时已构造,走 Drop 释放)。 +pub struct AzureAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(保留 api_key / resource_name / deployment_id 等原始字段) + config: ProviderConfig, + /// Azure 资源名称(如 "my-resource"),可为 None(当直接提供 base_url 时) + resource_name: Option, + /// Azure 部署 ID(作为 chat/image/embed 的路径段),可为 None(当直接提供 base_url 时) + deployment_id: Option, + /// Azure API 版本(如 "2024-02-15-preview") + api_version: String, + /// 用于 chat/image/embed 请求的 base_url(含 `/openai/deployments/{deployment}` 段) + /// + /// 当 `config.base_url` 提供时直接用之;否则由 `resource_name` + `deployment_id` + /// 拼接为 `https://{resource}.openai.azure.com/openai/deployments/{deployment}`。 + deployments_base_url: String, + /// OpenAI 兼容协议地基(复用请求体构造与响应解析,HTTP 层不使用它) + compat: OpenAiCompatAdapter, +} + +impl AzureAdapter { + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "azure"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Azure OpenAI"; + + /// 创建 Azure 适配器 + /// + /// 配置解析顺序(与 Python 老版 `__init__` 一致): + /// 1. `config.api_version` 为空时用 [`DEFAULT_AZURE_API_VERSION`] + /// 2. `config.base_url` 非空时直接作为 deployments base_url + /// 3. 否则用 `resource_name` + `deployment_id` 拼接默认 base_url + /// 4. 两者都没有则返 [`AibridgeError::Validation`](与 Python `ValueError` 一致) + /// + /// # 错误 + /// - `Validation`:既无 `base_url`,又无 `resource_name` + `deployment_id` + /// - HTTP 客户端构造失败(罕见,reqwest 配置错误) + pub fn new(config: ProviderConfig) -> Result { + let caps = Self::capabilities_set(); + + // 解析 api_version + let api_version = config + .api_version + .clone() + .filter(|v| !v.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_AZURE_API_VERSION.to_string()); + + // 解析 resource_name / deployment_id(从 config 直读,与 Python 一致) + let resource_name = config.resource_name.clone(); + let deployment_id = config.deployment_id.clone(); + + // 解析 deployments base_url + let deployments_base_url = if let Some(url) = config + .base_url + .as_ref() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + { + url + } else if let (Some(res), Some(dep)) = (resource_name.as_ref(), deployment_id.as_ref()) { + // 由 resource_name + deployment_id 拼接(去掉末尾斜杠,避免双斜杠) + let res = res.trim(); + let dep = dep.trim(); + format!("https://{res}.{DEFAULT_AZURE_HOST_SUFFIX}{AZURE_PATH_PREFIX}/{dep}") + } else { + return Err(AibridgeError::validation( + "Azure adapter requires either base_url or both resource_name and deployment_id", + )); + }; + + // 构造 HttpClient:用 deployments base_url 作为 base(chat/image/embed 路径相对它) + let http_opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(deployments_base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&http_opts)?; + + // 构造 OpenAiCompatAdapter 实例(仅复用其请求体构造/响应解析/错误映射方法, + // 其内部 HttpClient 不被本适配器使用)。传入相同 base_url 以保持一致性。 + let compat = OpenAiCompatAdapter::new( + config.clone(), + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + &deployments_base_url, + caps, + )?; + + Ok(Self { + http, + config, + resource_name, + deployment_id, + api_version, + deployments_base_url, + compat, + }) + } + + /// 支持的能力集合 + /// + /// 对应 Python v1 `AzureAdapter.supported_capabilities` 的核心子集 + /// (chat / chat_stream / image_generate / embedding)。 + /// Python 还声明了 AUDIO_TRANSCRIBE/TRANSLATE/SPEECH,但阶段 2a 范围 + /// 仅核心四能力,audio 走 trait 默认实现返 UnsupportedCapability。 + fn capabilities_set() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::ImageGenerate); + caps.insert(Capabilities::Embedding); + caps + } + + /// deployments base_url(含 `/openai/deployments/{deployment}` 段) + pub fn deployments_base_url(&self) -> &str { + &self.deployments_base_url + } + + /// Azure API 版本 + pub fn api_version(&self) -> &str { + &self.api_version + } + + /// Azure 部署 ID(chat/image/embed 路径段;直接提供 base_url 时为 None) + pub fn deployment_id(&self) -> Option<&str> { + self.deployment_id.as_deref() + } + + /// Azure 资源名称(用于 list_models 拼独立 host;直接提供 base_url 时为 None) + pub fn resource_name(&self) -> Option<&str> { + self.resource_name.as_deref() + } + + /// API key(可能为空) + fn api_key(&self) -> Option<&str> { + self.config.api_key.as_deref() + } + + /// 拼接完整 URL(deployments base_url + action 路径 + `?api-version=...`) + /// + /// `action` 如 `"chat/completions"` / `"images/generations"` / `"embeddings"`。 + /// 已含 query 的 URL 不再追加(如调用方自行拼好)。 + fn action_url(&self, action: &str) -> String { + let base = self.deployments_base_url.trim_end_matches('/'); + let action = action.trim_start_matches('/'); + format!("{base}/{action}?api-version={}", self.api_version) + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), Self::PROVIDER_TYPE), + }) + } + } + + /// 发送带 `api-key` header 的 POST JSON 请求,并用 OpenAI 错误映射处理响应 + /// + /// Azure 认证用 `api-key` header(非 Bearer),错误体结构与 OpenAI 一致 + /// (`{"error": {"message": "..."}}`),故复用 [`OpenAiCompatAdapter::map_api_error`]。 + async fn post_azure_json(&self, url: &str, body: &Value) -> Result { + let resp = self + .http + .inner() + .post(url) + .header("api-key", self.api_key().unwrap_or("")) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 `api-key` header 的 GET 请求,并用 OpenAI 错误映射处理响应 + async fn get_azure_json(&self, url: &str) -> Result { + let resp = self + .http + .inner() + .get(url) + .header("api-key", self.api_key().unwrap_or("")) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } +} + +#[async_trait] +impl Adapter for AzureAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + // 返回新集合,避免外部修改内部状态(不可变原则) + Self::capabilities_set() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new() 时已构造,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 走 Drop 释放,无需显式关闭 + Ok(()) + } + + /// 文本对话 + /// + /// `POST /openai/deployments/{deployment}/chat/completions?api-version={version}` + /// 请求体与 OpenAI 一致(委托地基构造),响应解析也委托地基。 + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = self.compat.build_chat_body(&req, false); + let url = self.action_url("chat/completions"); + let value = self.post_azure_json(&url, &body).await?; + self.compat.parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话 + /// + /// `POST /openai/deployments/{deployment}/chat/completions?api-version={version}&stream=true` + /// SSE 格式与 OpenAI 一致,按行解析 `data: ` / `data: [DONE]`。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = self.compat.build_chat_body(&req, true); + let url = self.action_url("chat/completions"); + + let resp = self + .http + .inner() + .post(&url) + .header("api-key", self.api_key().unwrap_or("")) + .json(&body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + + let model = req.model.clone(); + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + // 复用 openai_compat 的 LinesStream(按行切分字节流) + // 由于 LinesStream 是 openai_compat 的私有结构,这里内联同样的按行切分逻辑 + let lines_stream = AzureLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + if line.is_empty() || line.starts_with(':') { + continue; + } + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + if data.trim() == "[DONE]" { + return; + } + match serde_json::from_str::(data) { + Ok(v) => { + // 复用地基的 chunk 解析逻辑(parse_chunk 是关联函数,无需借 self) + match OpenAiCompatAdapter::parse_chunk(&v, &model) { + Ok(Some(chunk)) => yield Ok(chunk), + Ok(None) => continue, + Err(e) => { + yield Err(e); + return; + } + } + } + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 图像生成 + /// + /// `POST /openai/deployments/{deployment}/images/generations?api-version={version}` + async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let body = self.compat.build_image_body(&req); + let url = self.action_url("images/generations"); + let value = self.post_azure_json(&url, &body).await?; + self.compat.parse_image_result(&value, &req.model) + } + + /// 文本嵌入 + /// + /// `POST /openai/deployments/{deployment}/embeddings?api-version={version}` + async fn embed(&self, req: EmbedRequest) -> Result { + self.ensure_capability(Capabilities::Embedding)?; + let body = self.compat.build_embed_body(&req); + let url = self.action_url("embeddings"); + let value = self.post_azure_json(&url, &body).await?; + self.compat.parse_embedding_result(&value, &req.model) + } + + /// 模型列表(实时拉取 Azure 部署列表) + /// + /// `GET https://{resource}.openai.azure.com/openai/deployments?api-version={version}` + /// + /// Azure 部署列表响应:`{"data": [{"id": "<部署名>", "model": "<底层模型名>", ...}]}` + /// 与 Python 老版一致:用 `model` 字段作为模型 ID(用于推断类型),`id` 作为显示名。 + /// 若既无 `resource_name` 也无裸 base_url,则用 deployments_base_url 的 host 段拼路径。 + async fn list_models(&self, filter: Option) -> Result> { + // Azure 部署列表端点不含 deployments/{id} 段,路径为 /openai/deployments + // 取 resource host:优先用 resource_name 拼,否则从 deployments_base_url 反推 + let list_url = if let Some(res) = self.resource_name.as_ref() { + let res = res.trim(); + format!( + "https://{res}.{DEFAULT_AZURE_HOST_SUFFIX}{AZURE_PATH_PREFIX}?api-version={}", + self.api_version + ) + } else { + // 从 deployments_base_url(形如 .../openai/deployments/{dep})截到 /openai/deployments + let base = self.deployments_base_url.trim_end_matches('/'); + let prefix = AZURE_PATH_PREFIX; + let host_end = base.find(prefix).map(|i| &base[..i]).unwrap_or(base); + format!("{host_end}{prefix}?api-version={}", self.api_version) + }; + + let value = self.get_azure_json(&list_url).await?; + let models = parse_azure_deployments(&value, Self::PROVIDER_TYPE); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // 其余方法(video_create / video_poll / transcribe / speech / list_voices) + // 走 trait 默认实现,返 UnsupportedCapability。Python 老版的 transcribe/speech + // 待阶段 2c audio_adapters 统一补齐。 +} + +/// 解析 Azure 部署列表响应 → Vec +/// +/// Azure 响应:`{"data": [{"id": "<部署名>", "model": "<底层模型名>", ...}]}` +/// - 用 `model` 字段(缺失时回退 `id`)作为模型 ID,用于推断类型 +/// - 用 `id` 字段(部署名)作为显示名 +/// - `provider` 填 "azure" +fn parse_azure_deployments(value: &Value, provider: &str) -> Vec { + let arr = value.get("data").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .map(|m| { + // model 字段优先作为 ID(用于类型推断),缺失时回退 id + let model_field = m.get("model").and_then(|v| v.as_str()).unwrap_or(""); + let id_field = m.get("id").and_then(|v| v.as_str()).unwrap_or(""); + let model_id = if !model_field.is_empty() { + model_field + } else { + id_field + } + .to_string(); + let name = if !id_field.is_empty() { + id_field.to_string() + } else { + model_id.clone() + }; + let model_type = infer_model_type(&model_id); + ModelInfo { + name, + id: model_id, + model_type, + provider: provider.to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: None, + created: None, + } + }) + .collect(), + None => Vec::new(), + } +} + +// ==================== SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(与 openai_compat::LinesStream 等价实现) +/// +/// `openai_compat::LinesStream` 为私有结构,本适配器独立实现一份相同逻辑 +/// (按 `\n` 切分字节流,维护未完成行缓冲区)。 +struct AzureLinesStream { + inner: S, + buffer: Vec, +} + +impl AzureLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for AzureLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + } + Poll::Ready(None) => { + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::error::AibridgeError; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::options::{EmbedInput, EmbeddingVector}; + use futures::stream::StreamExt; + use mockito::Server; + use serde_json::json; + use std::collections::HashMap; + + /// 构造指向 mockito server 的 AzureAdapter(base_url 注入 mock 地址) + /// + /// mock 地址会被当作 deployments_base_url,故 chat 路径为 + /// `{mock}/chat/completions?api-version=...`。 + fn make_adapter(server: &Server) -> AzureAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("azure", opts); + AzureAdapter::new(config).expect("AzureAdapter 构造应成功") + } + + /// 构造 AzureAdapter(用 resource_name + deployment_id 拼默认 base_url,不发请求) + fn make_adapter_with_resource() -> AzureAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .extra("resource_name", "my-resource") + .extra("deployment_id", "my-deployment") + .extra("api_version", "2024-02-15-preview") + .build(); + let config = ProviderConfig::from_options("azure", opts); + AzureAdapter::new(config).expect("AzureAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 AzureAdapter(用于元信息/能力测试) + fn make_adapter_no_server() -> AzureAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url("https://example.openai.azure.com/openai/deployments/dep") + .build(); + let config = ProviderConfig::from_options("azure", opts); + AzureAdapter::new(config).expect("AzureAdapter 构造应成功") + } + + // ============ 构造与元信息 ============ + + #[test] + fn provider_type_and_name_match_python() { + let adapter = make_adapter_no_server(); + assert_eq!(adapter.provider_type(), "azure"); + assert_eq!(adapter.provider_name(), "Azure OpenAI"); + } + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_core_set() { + let adapter = make_adapter_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::ImageGenerate)); + assert!(caps.contains(&Capabilities::Embedding)); + // video / audio 不声明(阶段 2a 范围) + assert!(!caps.contains(&Capabilities::VideoGenerate)); + assert!(!caps.contains(&Capabilities::AudioSpeech)); + } + + #[test] + fn base_url_uses_config_when_provided() { + let adapter = make_adapter_no_server(); + assert_eq!( + adapter.deployments_base_url(), + "https://example.openai.azure.com/openai/deployments/dep" + ); + } + + #[test] + fn base_url_built_from_resource_and_deployment() { + let adapter = make_adapter_with_resource(); + assert_eq!( + adapter.deployments_base_url(), + "https://my-resource.openai.azure.com/openai/deployments/my-deployment" + ); + } + + #[test] + fn api_version_defaults_when_missing() { + // 未提供 api_version,应回退到默认值 + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://x.openai.azure.com/openai/deployments/d") + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + assert_eq!(adapter.api_version(), DEFAULT_AZURE_API_VERSION); + } + + #[test] + fn api_version_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://x.openai.azure.com/openai/deployments/d") + .extra("api_version", "2024-10-21") + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + assert_eq!(adapter.api_version(), "2024-10-21"); + } + + #[test] + fn new_fails_without_base_url_or_resource_deployment() { + // 既无 base_url,又无 resource_name + deployment_id → Validation 错误 + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("azure", opts); + let err = AzureAdapter::new(config) + .err() + .expect("应返回 Validation 错误"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn new_fails_with_only_resource_name() { + // 只有 resource_name,无 deployment_id → 失败 + let opts = ClientOptions::builder() + .api_key("k") + .extra("resource_name", "res") + .build(); + let config = ProviderConfig::from_options("azure", opts); + assert!(AzureAdapter::new(config).is_err()); + } + + #[test] + fn action_url_appends_api_version_query() { + let adapter = make_adapter_no_server(); + let url = adapter.action_url("chat/completions"); + assert!(url.contains("/chat/completions?api-version=")); + assert!(url.contains(DEFAULT_AZURE_API_VERSION)); + } + + // ============ chat 正常路径 ============ + + #[tokio::test] + async fn chat_success_returns_completion() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "gpt-4o", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }); + let mock = server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::UrlEncoded( + "api-version".into(), + DEFAULT_AZURE_API_VERSION.into(), + )) + .match_header("api-key", "test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.model, "gpt-4o"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_sends_api_key_header_not_bearer() { + // 关键:Azure 用 api-key header,而非 Authorization: Bearer + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .match_header("api-key", "test-key") + // 显式断言不带 Authorization Bearer(mockito 不匹配 Authorization 即可) + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gpt-4o", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_passes_extra_params_through() { + // Azure 请求体与 OpenAI 一致,extra 透传 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gpt-4o", + "custom_param": "custom_value" + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "gpt-4o", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]) + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + // ============ chat 错误路径 ============ + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(429) + .with_body( + json!({"error": {"message": "Rate limit exceeded", "retry_after": 1.5}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::RateLimit { retry_after, .. } => { + assert_eq!(retry_after, Some(1.5)); + } + _ => panic!("应为 RateLimit"), + } + } + + #[tokio::test] + async fn chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(404) + .with_body( + json!({"error": {"message": "The deployment 'gpt-x' does not exist"}}).to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-x", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(500) + .with_body(json!({"error": {"message": "Internal server error"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn chat_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(400) + .with_body(json!({"error": {"message": "max_tokens is invalid"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("max_tokens")); + } + _ => panic!("应为 Validation"), + } + } + + // ============ chat_stream ============ + + #[tokio::test] + async fn chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"gpt-4o\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .match_header("api-key", "test-key") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn chat_stream_sends_stream_true() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_query(mockito::Matcher::Any) + .match_body(mockito::Matcher::PartialJson(json!({ + "stream": true, + "stream_options": {"include_usage": true} + }))) + .with_status(200) + .with_body("data: [DONE]\n") + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + // ============ image_generate ============ + + #[tokio::test] + async fn image_generate_success_parses_url() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{ + "url": "https://example.com/img.png", + "revised_prompt": "a cute cat" + }] + }); + let mock = server + .mock("POST", "/images/generations") + .match_query(mockito::Matcher::UrlEncoded( + "api-version".into(), + DEFAULT_AZURE_API_VERSION.into(), + )) + .match_header("api-key", "test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat") + .size("1024x1024") + .build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://example.com/img.png") + ); + assert_eq!(resp.data[0].revised_prompt.as_deref(), Some("a cute cat")); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_success_parses_b64() { + let mut server = Server::new_async().await; + let body = json!({ + "created": 1700000000, + "data": [{"b64_json": "aGVsbG8="}] + }); + server + .mock("POST", "/images/generations") + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + } + + #[tokio::test] + async fn image_generate_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/images/generations") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server); + let req = ImageRequest::builder("dall-e-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ embed ============ + + #[tokio::test] + async fn embed_success_parses_vectors() { + let mut server = Server::new_async().await; + let body = json!({ + "object": "list", + "data": [ + {"object": "embedding", "index": 0, "embedding": [0.1, 0.2, 0.3]}, + {"object": "embedding", "index": 1, "embedding": [0.4, 0.5, 0.6]} + ], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 4, "total_tokens": 4} + }); + let mock = server + .mock("POST", "/embeddings") + .match_query(mockito::Matcher::UrlEncoded( + "api-version".into(), + DEFAULT_AZURE_API_VERSION.into(), + )) + .match_header("api-key", "test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Multiple(vec!["a".into(), "b".into()]), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 2); + assert_eq!(resp.data[0].index, 0); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2, 0.3]); + } else { + panic!("应为 Float 向量"); + } + assert_eq!(resp.usage.as_ref().unwrap().prompt_tokens, 4); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_single_input_sends_string() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/embeddings") + .match_query(mockito::Matcher::Any) + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "text-embedding-3-small", + "input": "hello" + }))) + .with_status(200) + .with_body( + json!({ + "object": "list", + "data": [{"object":"embedding","index":0,"embedding":[0.1]}], + "model": "text-embedding-3-small", + "usage": {"prompt_tokens": 1, "total_tokens": 1} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + mock.assert_async().await; + } + + #[tokio::test] + async fn embed_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_adapter(&server); + let req = EmbedRequest { + model: "text-embedding-3-small".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_success_parses_deployments() { + let mut server = Server::new_async().await; + // Azure 部署列表响应:id 是部署名,model 是底层模型名 + let body = json!({ + "data": [ + {"id": "my-gpt4o", "model": "gpt-4o", "object": "deployment", "status": "succeeded"}, + {"id": "my-dalle3", "model": "dall-e-3", "object": "deployment", "status": "succeeded"}, + {"id": "my-whisper", "model": "whisper-1", "object": "deployment", "status": "succeeded"} + ] + }); + // list_models 用 resource_name 拼独立 host,但这里用 base_url 路径反推 + // 让 server mock /openai/deployments 路径,base_url 设为 mock + /openai/deployments/dep + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(format!("{}/openai/deployments/dep", server.url())) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + + server + .mock("GET", "/openai/deployments") + .match_query(mockito::Matcher::UrlEncoded( + "api-version".into(), + DEFAULT_AZURE_API_VERSION.into(), + )) + .match_header("api-key", "test-key") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + // model 字段作为 ID(用于类型推断) + assert_eq!(models[0].id, "gpt-4o"); + assert_eq!(models[0].name, "my-gpt4o"); + assert_eq!(models[0].model_type, ModelType::Chat); + assert_eq!(models[0].provider, "azure"); + assert_eq!(models[1].model_type, ModelType::Image); + assert_eq!(models[2].model_type, ModelType::Audio); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "dep1", "model": "gpt-4o"}, + {"id": "dep2", "model": "dall-e-3"} + ] + }); + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(format!("{}/openai/deployments/dep", server.url())) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + + server + .mock("GET", "/openai/deployments") + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn list_models_falls_back_to_id_when_model_missing() { + // 部分 Azure 部署响应无 model 字段,回退用 id 作为模型 ID + let mut server = Server::new_async().await; + let body = json!({ + "data": [{"id": "custom-deploy"}] + }); + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(format!("{}/openai/deployments/dep", server.url())) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + + server + .mock("GET", "/openai/deployments") + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "custom-deploy"); + assert_eq!(models[0].name, "custom-deploy"); + } + + #[tokio::test] + async fn list_models_error_401_returns_authentication() { + let mut server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(format!("{}/openai/deployments/dep", server.url())) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = AzureAdapter::new(config).unwrap(); + + server + .mock("GET", "/openai/deployments") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ 不支持的能力走默认实现 ============ + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter_no_server(); + let req = crate::model::video::VideoRequest::builder("gpt-4o", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn speech_returns_unsupported() { + let adapter = make_adapter_no_server(); + let req = crate::model::audio::SpeechRequest::builder("tts-1", "hi", "alloy").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn list_voices_returns_unsupported() { + let adapter = make_adapter_no_server(); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ start / close ============ + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } +} diff --git a/crates/aibridge-core/src/adapters/openai_compat.rs b/crates/aibridge-core/src/adapters/openai_compat.rs index 4b53213..6552ac7 100644 --- a/crates/aibridge-core/src/adapters/openai_compat.rs +++ b/crates/aibridge-core/src/adapters/openai_compat.rs @@ -244,7 +244,10 @@ impl OpenAiCompatAdapter { /// 将统一 `ChatRequest` 转为 OpenAI 协议的 JSON 请求体。 /// 通用参数走 `OPENAI_COMPATIBLE_MAPPING`(当前为透传), /// provider 特有参数走 `extra` 透传。 - fn build_chat_body(&self, req: &ChatRequest, stream: bool) -> Value { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用请求体构造, + /// 仅 override HTTP path/header(见设计文档 10.1 节)。 + pub fn build_chat_body(&self, req: &ChatRequest, stream: bool) -> Value { // 序列化统一请求,得到基础字段 let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); // 强制覆盖 stream 标志(统一请求的 stream 字段默认 false,流式调用时需置 true) @@ -414,7 +417,9 @@ impl OpenAiCompatAdapter { // ==================== 内部:请求体构造 ==================== /// 构造 OpenAI images/generations 请求体 - fn build_image_body(&self, req: &ImageRequest) -> Value { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用请求体构造。 + pub fn build_image_body(&self, req: &ImageRequest) -> Value { let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); // extra 透传 if let Some(obj) = body.as_object_mut() { @@ -430,7 +435,9 @@ impl OpenAiCompatAdapter { } /// 构造 OpenAI embeddings 请求体 - fn build_embed_body(&self, req: &EmbedRequest) -> Value { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用请求体构造。 + pub fn build_embed_body(&self, req: &EmbedRequest) -> Value { let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); // extra 透传 if let Some(obj) = body.as_object_mut() { @@ -448,7 +455,13 @@ impl OpenAiCompatAdapter { // ==================== 内部:响应解析 ==================== /// 解析 OpenAI chat/completions 响应 → ChatCompletion - fn parse_chat_completion(&self, value: &Value, fallback_model: &str) -> Result { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用响应解析。 + pub fn parse_chat_completion( + &self, + value: &Value, + fallback_model: &str, + ) -> Result { let id = value .get("id") .and_then(|v| v.as_str()) @@ -504,7 +517,9 @@ impl OpenAiCompatAdapter { /// 解析单个 SSE chunk(OpenAI 流式格式)→ Option /// /// 返回 None 表示该 chunk 无有效 choices(如纯 usage 块),调用方跳过。 - fn parse_chunk(value: &Value, fallback_model: &str) -> Result> { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用流式 chunk 解析。 + pub fn parse_chunk(value: &Value, fallback_model: &str) -> Result> { let id = value .get("id") .and_then(|v| v.as_str()) @@ -552,7 +567,9 @@ impl OpenAiCompatAdapter { } /// 解析 OpenAI images/generations 响应 → ImageResult - fn parse_image_result(&self, value: &Value, fallback_model: &str) -> Result { + /// + /// 对子适配器开放(pub):Azure 等子适配器复用响应解析。 + pub fn parse_image_result(&self, value: &Value, fallback_model: &str) -> Result { let id = value .get("id") .and_then(|v| v.as_str()) @@ -587,7 +604,9 @@ impl OpenAiCompatAdapter { } /// 解析 OpenAI embeddings 响应 → EmbeddingResult - fn parse_embedding_result( + /// + /// 对子适配器开放(pub):Azure 等子适配器复用响应解析。 + pub fn parse_embedding_result( &self, value: &Value, fallback_model: &str, From 5fd3638627a701871273bfbb88d64bd0a87dba9c Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:42:27 +0800 Subject: [PATCH 23/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20aggregation=5Fplatforms=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/adapters/aggregation_platforms.rs | 1826 +++++++++++++++++ 1 file changed, 1826 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/aggregation_platforms.rs diff --git a/crates/aibridge-core/src/adapters/aggregation_platforms.rs b/crates/aibridge-core/src/adapters/aggregation_platforms.rs new file mode 100644 index 0000000..5ec28d4 --- /dev/null +++ b/crates/aibridge-core/src/adapters/aggregation_platforms.rs @@ -0,0 +1,1826 @@ +//! 聚合平台模型适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/aggregation_platforms.py`。 +//! +//! 支持四个聚合平台(均为 OpenAI 兼容协议族,差异在 base_url / 端点路径 / 响应结构): +//! - **SiliconFlow (硅基流动)**:标准 OpenAI 兼容,支持 chat/reasoning/embedding +//! - **Together AI**:标准 OpenAI 兼容,聚合开源模型 +//! - **Fireworks AI**:标准 OpenAI 兼容,高性能推理 +//! - **Cloudflare Workers AI**:边缘推理,使用 `/v1/run/{model}` 特殊端点 + `result` 响应结构 +//! +//! ## 结构 +//! +//! Python v1 是 4 个独立 adapter 类,统一注册到工厂。Rust 同样实现 4 个独立 struct: +//! - `SiliconFlowAdapter` / `TogetherAIAdapter` / `FireworksAIAdapter`:组合 `OpenAiCompatAdapter` +//! 地基,chat/chat_stream/image/embed/list_models 全部委托(仅 base_url/provider_type/capabilities +//! 差异) +//! - `CloudflareAIAdapter`:组合 `OpenAiCompatAdapter` 仅复用 embed(标准 /embeddings 端点), +//! chat/chat_stream/list_models 独立实现(Cloudflare 特有的 `/v1/run/{model}` 端点 + +//! `result` 响应结构 + `/models/search` 端点) +//! +//! ## provider_type 标识(对齐 Python v1) +//! +//! | 平台 | provider_type | 别名 | +//! |---|---|---| +//! | SiliconFlow | `siliconflow` | `sf` | +//! | Together AI | `togetherai` | `together` | +//! | Fireworks AI | `fireworksai` | `fireworks` | +//! | Cloudflare Workers AI | `cloudflareai` | `cloudflare` / `workersai` | + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatRequest, + ChoiceMessage, DeltaMessage, +}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::options::{EmbedRequest, EmbeddingResult}; +use crate::util; + +// ==================== 默认 Base URL ==================== + +/// SiliconFlow 默认 Base URL +/// +/// 对应 Python v1 `SiliconFlowAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_SILICONFLOW_BASE_URL: &str = "https://api.siliconflow.cn/v1"; + +/// Together AI 默认 Base URL +/// +/// 对应 Python v1 `TogetherAIAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_TOGETHERAI_BASE_URL: &str = "https://api.together.xyz/v1"; + +/// Fireworks AI 默认 Base URL +/// +/// 对应 Python v1 `FireworksAIAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_FIREWORKSAI_BASE_URL: &str = "https://api.fireworks.ai/inference/v1"; + +/// Cloudflare Workers AI 默认 Base URL(API 根,需拼接 account_id) +/// +/// 对应 Python v1 `CloudflareAIAdapter.DEFAULT_BASE_URL`。 +/// 实际请求 base 为 `{DEFAULT_CLOUDFLAREAI_BASE_URL}/accounts/{account_id}/ai`。 +pub const DEFAULT_CLOUDFLAREAI_BASE_URL: &str = "https://api.cloudflare.com/client/v4"; + +// ==================== 能力集合构造 ==================== + +/// SiliconFlow 支持的能力集合 +/// +/// 对齐 Python v1 `SiliconFlowAdapter.supported_capabilities`。 +fn siliconflow_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::ToolCall); + caps.insert(Capabilities::Reasoning); + caps.insert(Capabilities::JsonMode); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::AudioTranscribe); + caps.insert(Capabilities::AudioSpeech); + caps +} + +/// Together AI 支持的能力集合 +/// +/// 对齐 Python v1 `TogetherAIAdapter.supported_capabilities`。 +fn togetherai_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::JsonMode); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::AudioTranscribe); + caps.insert(Capabilities::AudioSpeech); + caps +} + +/// Fireworks AI 支持的能力集合 +/// +/// 对齐 Python v1 `FireworksAIAdapter.supported_capabilities`。 +fn fireworksai_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::ToolCall); + caps.insert(Capabilities::JsonMode); + caps.insert(Capabilities::Embedding); + caps.insert(Capabilities::AudioTranscribe); + caps +} + +/// Cloudflare Workers AI 支持的能力集合 +/// +/// 对齐 Python v1 `CloudflareAIAdapter.supported_capabilities`。 +fn cloudflareai_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::Embedding); + caps +} + +// ==================== SiliconFlow 适配器 ==================== + +/// SiliconFlow (硅基流动) 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.siliconflow.cn/v1` +/// - Chat: `POST /chat/completions`(支持 reasoning 等特有参数,经 `extra` 透传) +/// - Embed: `POST /embeddings` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct SiliconFlowAdapter { + /// OpenAI 兼容地基(chat/chat_stream/image/embed/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl SiliconFlowAdapter { + /// 创建 SiliconFlow 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "siliconflow", + "SiliconFlow 硅基流动", + DEFAULT_SILICONFLOW_BASE_URL, + siliconflow_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for SiliconFlowAdapter { + fn provider_type(&self) -> &str { + "siliconflow" + } + + fn provider_name(&self) -> &str { + "SiliconFlow 硅基流动" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Together AI 适配器 ==================== + +/// Together AI 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.together.xyz/v1` +/// - Chat: `POST /chat/completions` +/// - Embed: `POST /embeddings` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct TogetherAIAdapter { + /// OpenAI 兼容地基(chat/chat_stream/image/embed/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl TogetherAIAdapter { + /// 创建 Together AI 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "togetherai", + "Together AI", + DEFAULT_TOGETHERAI_BASE_URL, + togetherai_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for TogetherAIAdapter { + fn provider_type(&self) -> &str { + "togetherai" + } + + fn provider_name(&self) -> &str { + "Together AI" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Fireworks AI 适配器 ==================== + +/// Fireworks AI 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.fireworks.ai/inference/v1` +/// - Chat: `POST /chat/completions` +/// - Embed: `POST /embeddings` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct FireworksAIAdapter { + /// OpenAI 兼容地基(chat/chat_stream/image/embed/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl FireworksAIAdapter { + /// 创建 Fireworks AI 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "fireworksai", + "Fireworks AI", + DEFAULT_FIREWORKSAI_BASE_URL, + fireworksai_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for FireworksAIAdapter { + fn provider_type(&self) -> &str { + "fireworksai" + } + + fn provider_name(&self) -> &str { + "Fireworks AI" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Cloudflare Workers AI 适配器 ==================== + +/// Cloudflare Workers AI 适配器 +/// +/// 部分兼容 OpenAI 协议,但 chat 端点与响应结构与标准 OpenAI 不同: +/// - Base URL: `https://api.cloudflare.com/client/v4/accounts/{account_id}/ai` +/// - Chat: `POST /v1/run/{model}`(模型用 `@cf/` 前缀) +/// - 响应: `{"result": {"response": "..."}}`(非标准 OpenAI choices 结构) +/// - Models: `GET /models/search`(响应 `{"result": {"models": [...]}}`) +/// - Embed: `POST /embeddings`(标准 OpenAI 兼容,委托 compat) +/// - 认证: Bearer Token + Account ID(account_id 从 `config.extra["account_id"]` 取) +/// +/// `account_id` 为必填,构造时缺失则 `new()` 返回 `Validation` 错误。 +pub struct CloudflareAIAdapter { + /// OpenAI 兼容地基(仅 embed 委托给它,标准 /embeddings 端点) + compat: OpenAiCompatAdapter, + /// chat/list_models 端点专用 HTTP 客户端(独立于 compat,避免暴露 compat 私有字段) + http: HttpClient, + /// Cloudflare account_id(构造时校验非空,用于构造 ai_base;保留供调试/扩展) + #[allow(dead_code)] + account_id: String, + /// Provider 配置(保留引用以便取 api_key) + config: ProviderConfig, +} + +impl CloudflareAIAdapter { + /// 创建 Cloudflare 适配器 + /// + /// - `config.extra["account_id"]` 必须存在且非空,否则返回 `Validation` 错误 + /// - `config.base_url` 为 None 时用 `DEFAULT_CLOUDFLAREAI_BASE_URL` 兜底 + pub fn new(config: ProviderConfig) -> Result { + let account_id = config + .extra + .get("account_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .filter(|s| !s.trim().is_empty()) + .ok_or_else(|| { + AibridgeError::validation( + "Cloudflare account_id 不能为空(请配置 extra.account_id)", + ) + })?; + + let api_root = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_CLOUDFLAREAI_BASE_URL.to_string()); + + // Cloudflare 实际请求 base:{api_root}/accounts/{account_id}/ai + let ai_base = format!( + "{}/accounts/{}/ai", + api_root.trim_end_matches('/'), + account_id + ); + + // compat 用于 embed(标准 /embeddings 端点),需让其 base_url 指向 ai_base。 + // OpenAiCompatAdapter::new 会优先用 config.base_url,故这里把 config 副本的 + // base_url 改写为 ai_base 后再传入,确保 embed 端点走 /accounts/{id}/ai/embeddings。 + let mut compat_config = config.clone(); + compat_config.base_url = Some(ai_base.clone()); + let compat = OpenAiCompatAdapter::new( + compat_config, + "cloudflareai", + "Cloudflare Workers AI", + &ai_base, + cloudflareai_capabilities(), + )?; + + // chat/list_models 端点专用 HttpClient,base_url 同 ai_base + let http = HttpClient::new( + &ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(ai_base) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(), + )?; + + Ok(Self { + compat, + http, + account_id, + config, + }) + } + + /// 用显式 HttpClient + compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat( + compat: OpenAiCompatAdapter, + http: HttpClient, + account_id: impl Into, + config: ProviderConfig, + ) -> Self { + Self { + compat, + http, + account_id: account_id.into(), + config, + } + } + + /// account_id(测试可见) + #[cfg(test)] + pub fn account_id(&self) -> &str { + &self.account_id + } + + /// API key + fn api_key(&self) -> Option<&str> { + self.config.api_key.as_deref() + } + + /// base_url(compat 持有的 ai_base) + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持(Cloudflare 独立实现,因 ensure_capability 在 compat 内为私有) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: cloudflareai)", cap.as_str()), + }) + } + } + + /// 发送带认证的 POST JSON 请求,并用 OpenAI 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带认证的 GET 请求,并用 OpenAI 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 格式化 Cloudflare 模型名(补充 `@cf/` 前缀) + /// + /// 对应 Python v1 `CloudflareAIAdapter._format_model`。 + fn format_model(model: &str) -> String { + if model.starts_with("@cf/") { + model.to_string() + } else { + format!("@cf/{model}") + } + } + + /// 构造 Cloudflare chat 请求体 + /// + /// 对应 Python v1 `CloudflareAIAdapter.chat` 的 body 构造: + /// - 不含 `model`(模型在 URL 路径里) + /// - 含 `messages` / `temperature` / `max_tokens` / `stream` + /// - `extra` 透传 + fn build_chat_body(req: &ChatRequest, stream: bool) -> Value { + let mut body = serde_json::to_value(req).unwrap_or_else(|_| json!({})); + if let Some(obj) = body.as_object_mut() { + // 移除 model(Cloudflare 模型在 URL 路径,不在 body) + obj.remove("model"); + if stream { + obj.insert("stream".to_string(), json!(true)); + } else if obj.get("stream").and_then(|v| v.as_bool()).unwrap_or(false) { + obj.insert("stream".to_string(), json!(false)); + } + // extra 透传到顶层 + if let Some(extra) = obj.remove("extra") { + if let Some(extra_map) = extra.as_object() { + for (k, v) in extra_map { + obj.insert(k.clone(), v.clone()); + } + } + } + } + body + } + + /// 解析 Cloudflare chat 响应 → ChatCompletion + /// + /// Cloudflare 响应格式:`{"result": {"response": "..."}}`(非标准 OpenAI choices)。 + /// 对应 Python v1 `CloudflareAIAdapter._parse_response`。 + fn parse_chat_completion(value: &Value, fallback_model: &str) -> Result { + let result = value.get("result").unwrap_or(&Value::Null); + let content = if result.is_object() { + result + .get("response") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string() + } else { + // result 不是对象时,转为字符串作为 content + result.as_str().unwrap_or("").to_string() + }; + + let usage = value.get("usage").and_then(parse_usage); + + Ok(ChatCompletion { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion".to_string(), + created: value + .get("timestamp") + .and_then(|v| v.as_u64()) + .or_else(|| value.get("created").and_then(|v| v.as_u64())) + .unwrap_or_else(util::current_timestamp), + model: fallback_model.to_string(), + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".to_string(), + content: Some(content), + tool_calls: None, + }, + finish_reason: Some("stop".to_string()), + }], + usage, + service_tier: None, + system_fingerprint: None, + }) + } + + /// 解析单个 Cloudflare 流式 chunk + /// + /// Cloudflare 流式响应:每个 SSE data 是 `{"result": {"response": "增量文本"}}`。 + /// 对应 Python v1 `CloudflareAIAdapter._parse_chunk`。 + fn parse_chunk(value: &Value, fallback_model: &str) -> Option { + let result = value.get("result").unwrap_or(&Value::Null); + let content = if result.is_object() { + result + .get("response") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string() + } else { + String::new() + }; + + // 空内容块跳过(与 Python 老版一致:无有效内容返回 None) + if content.is_empty() { + return None; + } + + Some(ChatCompletionChunk { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion.chunk".to_string(), + created: util::current_timestamp(), + model: fallback_model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".to_string()), + content: Some(content), + tool_calls: None, + }, + finish_reason: None, + }], + usage: None, + }) + } + + /// 解析 Cloudflare /models/search 响应 → Vec + /// + /// Cloudflare 响应:`{"result": {"models": [...]}}`,兼容 result 为 list 的情况。 + /// 对应 Python v1 `CloudflareAIAdapter.list_models`。 + fn parse_models(value: &Value, provider: &str) -> Vec { + let result = value.get("result").unwrap_or(&Value::Null); + let arr: Option<&Vec> = if result.is_object() { + result.get("models").and_then(|v| v.as_array()) + } else { + // result 直接是 list 的情况 + result.as_array() + }; + + match arr { + Some(arr) => arr + .iter() + .map(|m| { + // Cloudflare 模型条目可能是字符串或对象 {"id": "...", "name": "..."} + let id = if let Some(s) = m.as_str() { + s.to_string() + } else { + m.get("id") + .or_else(|| m.get("name")) + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string() + }; + let model_type = infer_model_type(&id); + ModelInfo { + name: id.clone(), + id, + model_type, + provider: provider.to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: m + .get("description") + .and_then(|v| v.as_str()) + .map(str::to_owned), + created: None, + } + }) + .collect(), + None => Vec::new(), + } + } +} + +#[async_trait] +impl Adapter for CloudflareAIAdapter { + fn provider_type(&self) -> &str { + "cloudflareai" + } + + fn provider_name(&self) -> &str { + "Cloudflare Workers AI" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话(Cloudflare 特有协议:POST /v1/run/{model}) + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let cf_model = Self::format_model(&req.model); + let body = Self::build_chat_body(&req, false); + let value = self + .post_authed_json(&format!("v1/run/{cf_model}"), &body) + .await?; + Self::parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话(Cloudflare 特有协议:POST /v1/run/{model} stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let cf_model = Self::format_model(&req.model); + let body = Self::build_chat_body(&req, true); + let url = self.url(&format!("v1/run/{cf_model}")); + + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(&body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + + let model = req.model.clone(); + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = CloudflareLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + if line.is_empty() || line.starts_with(':') { + continue; + } + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + if data.trim() == "[DONE]" { + return; + } + match serde_json::from_str::(data) { + Ok(v) => match Self::parse_chunk(&v, &model) { + Some(chunk) => yield Ok(chunk), + None => continue, + }, + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 文本嵌入(委托给 OpenAiCompatAdapter,标准 /embeddings 端点) + /// + /// Cloudflare Workers AI 同样提供标准 OpenAI 兼容 embeddings 端点, + /// 复用地基实现。Python v1 声明了 EMBEDDING 能力但未实现 embed(默认返不支持), + /// 此处补齐为标准 OpenAI 兼容实现。 + async fn embed(&self, req: EmbedRequest) -> Result { + self.compat.embed(req).await + } + + /// 模型列表(Cloudflare 特有:GET /models/search) + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("models/search").await?; + let models = Self::parse_models(&value, self.provider_type()); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +/// 解析 usage 统计 +/// +/// 与 `openai_compat::parse_usage` 同结构,本模块独立持有副本(原函数为模块私有)。 +fn parse_usage(v: &Value) -> Option { + let prompt = v.get("prompt_tokens").and_then(|x| x.as_u64())?; + let completion = v + .get("completion_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(0); + let total = v + .get("total_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(prompt + completion); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: total, + }) +} + +// ==================== Cloudflare SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(Cloudflare 流式用) +/// +/// 与 `openai_compat::LinesStream` 等价,独立实现避免引用其私有结构。 +struct CloudflareLinesStream { + inner: S, + buffer: Vec, +} + +impl CloudflareLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for CloudflareLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + } + Poll::Ready(None) => { + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::http::HttpClient; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::image::ImageRequest; + use crate::model::options::{EmbedInput, EmbeddingVector}; + use crate::model::video::VideoRequest; + use mockito::Server; + use std::collections::HashMap; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 OpenAiCompatAdapter(指向 mockito server,给定 provider 信息与能力) + fn make_compat( + server: &Server, + provider_type: &str, + provider_name: &str, + caps: CapabilitySet, + ) -> OpenAiCompatAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options(provider_type, opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + OpenAiCompatAdapter::with_http(http, config, provider_type, provider_name, caps) + } + + /// 构造测试用 SiliconFlowAdapter(指向 mockito server) + fn make_siliconflow(server: &Server) -> SiliconFlowAdapter { + let compat = make_compat( + server, + "siliconflow", + "SiliconFlow 硅基流动", + siliconflow_capabilities(), + ); + SiliconFlowAdapter::with_compat(compat) + } + + /// 构造测试用 TogetherAIAdapter(指向 mockito server) + fn make_togetherai(server: &Server) -> TogetherAIAdapter { + let compat = make_compat( + server, + "togetherai", + "Together AI", + togetherai_capabilities(), + ); + TogetherAIAdapter::with_compat(compat) + } + + /// 构造测试用 FireworksAIAdapter(指向 mockito server) + fn make_fireworksai(server: &Server) -> FireworksAIAdapter { + let compat = make_compat( + server, + "fireworksai", + "Fireworks AI", + fireworksai_capabilities(), + ); + FireworksAIAdapter::with_compat(compat) + } + + /// 构造测试用 CloudflareAIAdapter(指向 mockito server) + /// + /// 需注入 account_id 到 config.extra,并让 compat/chat http 的 base_url 都指向 + /// `{server}/accounts/test-account/ai`(模拟 Cloudflare ai_base)。 + /// 注意:config.base_url 必须设为 ai_base,因为 OpenAiCompatAdapter::with_http + /// 用 config.base_url 作为其 base_url(而非传入 http 的 base_url)。 + fn make_cloudflare(server: &Server) -> CloudflareAIAdapter { + let ai_base = format!( + "{}/accounts/test-account/ai", + server.url().trim_end_matches('/') + ); + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(ai_base.clone()) + .timeout(5) + .extra("account_id", json!("test-account")) + .build(); + let config = ProviderConfig::from_options("cloudflareai", opts); + let compat_http = + HttpClient::new(&ClientOptions::builder().base_url(&ai_base).build()).unwrap(); + let compat = OpenAiCompatAdapter::with_http( + compat_http, + config.clone(), + "cloudflareai", + "Cloudflare Workers AI", + cloudflareai_capabilities(), + ); + let chat_http = + HttpClient::new(&ClientOptions::builder().base_url(&ai_base).build()).unwrap(); + CloudflareAIAdapter::with_compat(compat, chat_http, "test-account", config) + } + + /// 标准 OpenAI chat 成功响应体 + fn openai_chat_body() -> Value { + json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "test-model", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }) + } + + // ==================== SiliconFlow 测试 ==================== + + #[tokio::test] + async fn siliconflow_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_siliconflow(&server); + assert_eq!(adapter.provider_type(), "siliconflow"); + assert_eq!(adapter.provider_name(), "SiliconFlow 硅基流动"); + } + + #[tokio::test] + async fn siliconflow_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_siliconflow(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn siliconflow_capabilities_include_chat_embed_reasoning() { + let server = Server::new_async().await; + let adapter = make_siliconflow(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Reasoning)); + assert!(caps.contains(&Capabilities::Embedding)); + assert!(caps.contains(&Capabilities::JsonMode)); + } + + #[tokio::test] + async fn siliconflow_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let req = ChatRequest::builder("Qwen/Qwen2.5-7B-Instruct", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn siliconflow_chat_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "invalid key"}}).to_string()) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn siliconflow_embed_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(200) + .with_body( + json!({ + "object": "list", + "data": [{"object":"embedding","index":0,"embedding":[0.1,0.2]}], + "model": "BAAI/bge-large-zh-v1.5", + "usage": {"prompt_tokens": 2, "total_tokens": 2} + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let req = EmbedRequest { + model: "BAAI/bge-large-zh-v1.5".into(), + input: EmbedInput::Single("hello".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.1, 0.2]); + } else { + panic!("应为 Float 向量"); + } + } + + #[tokio::test] + async fn siliconflow_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "Qwen/Qwen2.5-7B-Instruct", "object": "model"}, + {"id": "BAAI/bge-large-zh-v1.5", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "Qwen/Qwen2.5-7B-Instruct"); + assert_eq!(models[0].provider, "siliconflow"); + } + + #[tokio::test] + async fn siliconflow_list_models_filter_by_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "Qwen/Qwen2.5-7B-Instruct", "object": "model"}, + {"id": "dall-e-3", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn siliconflow_list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_siliconflow(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ==================== Together AI 测试 ==================== + + #[tokio::test] + async fn togetherai_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_togetherai(&server); + assert_eq!(adapter.provider_type(), "togetherai"); + assert_eq!(adapter.provider_name(), "Together AI"); + } + + #[tokio::test] + async fn togetherai_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "meta-llama/Llama-3-70B-chat-hf", + "temperature": 0.5 + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_togetherai(&server); + let req = ChatRequest::builder( + "meta-llama/Llama-3-70B-chat-hf", + vec![ChatMessage::user("hi")], + ) + .temperature(0.5) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn togetherai_chat_error_429() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "rate limit"}}).to_string()) + .create_async() + .await; + let adapter = make_togetherai(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn togetherai_embed_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(200) + .with_body( + json!({ + "data": [{"object":"embedding","index":0,"embedding":[0.9]}], + "model": "BAAI/bge-base-en-v1.5" + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_togetherai(&server); + let req = EmbedRequest { + model: "BAAI/bge-base-en-v1.5".into(), + input: EmbedInput::Single("text".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + } + + #[tokio::test] + async fn togetherai_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "meta-llama/Llama-3-70B-chat-hf"}, + {"id": "mistralai/Mixtral-8x7B-Instruct-v0.1"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_togetherai(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].provider, "togetherai"); + } + + // ==================== Fireworks AI 测试 ==================== + + #[tokio::test] + async fn fireworksai_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_fireworksai(&server); + assert_eq!(adapter.provider_type(), "fireworksai"); + assert_eq!(adapter.provider_name(), "Fireworks AI"); + } + + #[tokio::test] + async fn fireworksai_capabilities_include_tool_call() { + let server = Server::new_async().await; + let adapter = make_fireworksai(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::ToolCall)); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::Embedding)); + assert!(caps.contains(&Capabilities::AudioTranscribe)); + } + + #[tokio::test] + async fn fireworksai_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "accounts/fireworks/models/llama-v3-70b-instruct" + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_fireworksai(&server); + let req = ChatRequest::builder( + "accounts/fireworks/models/llama-v3-70b-instruct", + vec![ChatMessage::user("hi")], + ) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn fireworksai_chat_error_404() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body(json!({"error": {"message": "model not found"}}).to_string()) + .create_async() + .await; + let adapter = make_fireworksai(&server); + let req = ChatRequest::builder("bad-model", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn fireworksai_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [{"id": "accounts/fireworks/models/llama-v3-70b-instruct"}] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_fireworksai(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].provider, "fireworksai"); + } + + #[tokio::test] + async fn fireworksai_embed_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/embeddings") + .with_status(200) + .with_body( + json!({ + "data": [{"object":"embedding","index":0,"embedding":[1.0,2.0]}], + "model": "nomic-ai/nomic-embed-text-v1.5" + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_fireworksai(&server); + let req = EmbedRequest { + model: "nomic-ai/nomic-embed-text-v1.5".into(), + input: EmbedInput::Single("text".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + } + + // ==================== Cloudflare 测试 ==================== + + #[tokio::test] + async fn cloudflare_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_cloudflare(&server); + assert_eq!(adapter.provider_type(), "cloudflareai"); + assert_eq!(adapter.provider_name(), "Cloudflare Workers AI"); + } + + #[tokio::test] + async fn cloudflare_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_cloudflare(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn cloudflare_account_id_injected() { + let server = Server::new_async().await; + let adapter = make_cloudflare(&server); + assert_eq!(adapter.account_id(), "test-account"); + } + + #[tokio::test] + async fn cloudflare_new_without_account_id_returns_validation_error() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://example.com") + .build(); + let config = ProviderConfig::from_options("cloudflareai", opts); + let result = CloudflareAIAdapter::new(config); + assert!(matches!(result, Err(AibridgeError::Validation { .. }))); + } + + #[tokio::test] + async fn cloudflare_format_model_adds_cf_prefix() { + assert_eq!( + CloudflareAIAdapter::format_model("llama-3-8b"), + "@cf/llama-3-8b" + ); + assert_eq!( + CloudflareAIAdapter::format_model("@cf/llama-3-8b"), + "@cf/llama-3-8b" + ); + } + + #[tokio::test] + async fn cloudflare_chat_success() { + let mut server = Server::new_async().await; + // Cloudflare 端点:/accounts/test-account/ai/v1/run/@cf/meta/llama-3-8b-instruct + let mock = server + .mock( + "POST", + "/accounts/test-account/ai/v1/run/@cf/meta/llama-3-8b-instruct", + ) + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "result": {"response": "Hi from Cloudflare!"}, + "usage": {"prompt_tokens": 3, "completion_tokens": 4, "total_tokens": 7} + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = ChatRequest::builder( + "@cf/meta/llama-3-8b-instruct", + vec![ChatMessage::user("hi")], + ) + .temperature(0.5) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.choices.len(), 1); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("Hi from Cloudflare!") + ); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn cloudflare_chat_auto_prefixes_cf() { + // 不带 @cf/ 前缀的模型名应被自动补全 + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/accounts/test-account/ai/v1/run/@cf/llama-3-8b-instruct", + ) + .with_status(200) + .with_body(json!({"result": {"response": "ok"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = + ChatRequest::builder("llama-3-8b-instruct", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("ok")); + mock.assert_async().await; + } + + #[tokio::test] + async fn cloudflare_chat_body_excludes_model_and_keeps_params() { + // 验证请求体不含 model 字段,但保留 messages / temperature + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/accounts/test-account/ai/v1/run/@cf/test-model") + .match_body(mockito::Matcher::PartialJson(json!({ + "messages": [{"role": "user", "content": "hi"}], + "temperature": 0.7 + }))) + .with_status(200) + .with_body(json!({"result": {"response": "ok"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = ChatRequest::builder("test-model", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("ok")); + mock.assert_async().await; + } + + #[tokio::test] + async fn cloudflare_chat_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/accounts/test-account/ai/v1/run/@cf/m") + .with_status(401) + .with_body(json!({"error": {"message": "bad token"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn cloudflare_chat_error_429() { + let mut server = Server::new_async().await; + server + .mock("POST", "/accounts/test-account/ai/v1/run/@cf/m") + .with_status(429) + .with_body(json!({"error": {"message": "slow"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn cloudflare_chat_stream_success() { + let mut server = Server::new_async().await; + // SSE 流:两个增量块 + [DONE] + let sse = "data: {\"result\":{\"response\":\"Hello\"}}\n\ + data: {\"result\":{\"response\":\" world\"}}\n\ + data: [DONE]\n"; + let mock = server + .mock("POST", "/accounts/test-account/ai/v1/run/@cf/m") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let stream = adapter.chat_stream(req).await.expect("流应建立"); + let chunks: Vec<_> = stream.collect().await; + let contents: Vec = chunks + .into_iter() + .filter_map(|c| c.ok()) + .filter_map(|c| c.choices.into_iter().next()) + .filter_map(|d| d.delta.content) + .collect(); + assert_eq!(contents, vec!["Hello", " world"]); + mock.assert_async().await; + } + + #[tokio::test] + async fn cloudflare_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/accounts/test-account/ai/models/search") + .with_status(200) + .with_body( + json!({ + "result": { + "models": [ + {"id": "@cf/meta/llama-3-8b-instruct", "description": "Llama 3 8B"}, + {"id": "@cf/baai/bge-base-en-v1.5", "description": "Embedding"} + ] + } + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "@cf/meta/llama-3-8b-instruct"); + assert_eq!(models[0].provider, "cloudflareai"); + assert_eq!(models[0].description.as_deref(), Some("Llama 3 8B")); + } + + #[tokio::test] + async fn cloudflare_list_models_result_as_list() { + // 兼容 result 直接为 list 的情况 + let mut server = Server::new_async().await; + server + .mock("GET", "/accounts/test-account/ai/models/search") + .with_status(200) + .with_body( + json!({ + "result": [{"id": "@cf/m1"}, "@cf/m2"] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "@cf/m1"); + assert_eq!(models[1].id, "@cf/m2"); + } + + #[tokio::test] + async fn cloudflare_list_models_filter_by_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/accounts/test-account/ai/models/search") + .with_status(200) + .with_body( + json!({ + "result": { + "models": [ + {"id": "@cf/meta/llama-3-8b-instruct"}, + {"id": "@cf/stabilityai/stable-diffusion-xl-base-1.0"} + ] + } + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "@cf/stabilityai/stable-diffusion-xl-base-1.0"); + } + + #[tokio::test] + async fn cloudflare_list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/accounts/test-account/ai/models/search") + .with_status(401) + .with_body(json!({"error": {"message": "bad token"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn cloudflare_embed_success() { + // embed 委托 compat(标准 /embeddings 端点) + let mut server = Server::new_async().await; + server + .mock("POST", "/accounts/test-account/ai/embeddings") + .with_status(200) + .with_body( + json!({ + "data": [{"object":"embedding","index":0,"embedding":[0.5,0.6]}], + "model": "@cf/baai/bge-base-en-v1.5" + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = EmbedRequest { + model: "@cf/baai/bge-base-en-v1.5".into(), + input: EmbedInput::Single("text".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let resp = adapter.embed(req).await.unwrap(); + assert_eq!(resp.data.len(), 1); + if let EmbeddingVector::Float(v) = &resp.data[0].embedding { + assert_eq!(v, &vec![0.5, 0.6]); + } else { + panic!("应为 Float 向量"); + } + } + + #[tokio::test] + async fn cloudflare_embed_error_401() { + let mut server = Server::new_async().await; + server + .mock("POST", "/accounts/test-account/ai/embeddings") + .with_status(401) + .with_body(json!({"error": {"message": "bad token"}}).to_string()) + .create_async() + .await; + let adapter = make_cloudflare(&server); + let req = EmbedRequest { + model: "@cf/baai/bge-base-en-v1.5".into(), + input: EmbedInput::Single("text".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ==================== 能力校验(不支持的能力返回错误) ==================== + + #[tokio::test] + async fn cloudflare_image_generate_unsupported() { + let server = Server::new_async().await; + let adapter = make_cloudflare(&server); + let req = ImageRequest::builder("m", "prompt").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn siliconflow_video_create_unsupported() { + let server = Server::new_async().await; + let adapter = make_siliconflow(&server); + let req = VideoRequest::builder("m", "p").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== 默认 Base URL 常量校验 ==================== + + #[test] + fn default_base_urls_match_python() { + assert_eq!( + DEFAULT_SILICONFLOW_BASE_URL, + "https://api.siliconflow.cn/v1" + ); + assert_eq!(DEFAULT_TOGETHERAI_BASE_URL, "https://api.together.xyz/v1"); + assert_eq!( + DEFAULT_FIREWORKSAI_BASE_URL, + "https://api.fireworks.ai/inference/v1" + ); + assert_eq!( + DEFAULT_CLOUDFLAREAI_BASE_URL, + "https://api.cloudflare.com/client/v4" + ); + } + + // ==================== start/close 无副作用 ==================== + + #[tokio::test] + async fn siliconflow_start_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_siliconflow(&server); + adapter.start().await.unwrap(); + adapter.close().await.unwrap(); + } + + #[tokio::test] + async fn cloudflare_start_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_cloudflare(&server); + adapter.start().await.unwrap(); + adapter.close().await.unwrap(); + } +} From cb4536a8c516fa433ee2dc1a39dfe03af842580b Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:50:46 +0800 Subject: [PATCH 24/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20=E7=AC=AC=E4=B8=80=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20azure=20+=20=E8=81=9A=E5=90=88=E5=B9=B3=E5=8F=B0?= =?UTF-8?q?=E5=88=B0=E5=B7=A5=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 118 +++++++++++++++++--- crates/aibridge-core/src/adapters/gemini.rs | 2 +- crates/aibridge-core/src/adapters/mod.rs | 6 + 3 files changed, 108 insertions(+), 18 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index 4c7e976..8b0c89d 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -8,7 +8,11 @@ //! - 阶段 0.4 暂只占位分支(返 ProviderNotFound),具体适配器阶段 1 起填充 use crate::adapter::Adapter; +use crate::adapters::aggregation_platforms::{ + CloudflareAIAdapter, FireworksAIAdapter, SiliconFlowAdapter, TogetherAIAdapter, +}; use crate::adapters::agnes::AgnesAdapter; +use crate::adapters::azure::AzureAdapter; use crate::adapters::echo::EchoAdapter; use crate::adapters::gemini::GeminiAdapter; use crate::adapters::openai::OpenAiAdapter; @@ -26,15 +30,19 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "agnes", "volcengine_cv", "gemini", - // 阶段 2 起补充: + // 阶段 2a 已实现(兼容族): "azure", + "siliconflow", + "togetherai", + "fireworksai", + "cloudflareai", + // 阶段 2b/2c 待实现: "anthropic", "runway", "pika", "kling", "stability", "chinese", - "aggregation_platforms", "edge-tts", "elevenlabs", "cartesia", @@ -57,22 +65,22 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { "agnes" => Ok(Box::new(AgnesAdapter::new(config)?)), "volcengine_cv" => Ok(Box::new(VolcengineCvAdapter::new(config)?)), "gemini" => Ok(Box::new(GeminiAdapter::new(config)?)), + // 阶段 2a 适配器:OpenAI 兼容族 + "azure" => Ok(Box::new(AzureAdapter::new(config)?)), + // 聚合平台:别名对齐 Python agn/adapters/aggregation_platforms.py 末尾 register 调用 + "siliconflow" | "sf" => Ok(Box::new(SiliconFlowAdapter::new(config)?)), + "togetherai" | "together" => Ok(Box::new(TogetherAIAdapter::new(config)?)), + "fireworksai" | "fireworks" => Ok(Box::new(FireworksAIAdapter::new(config)?)), + "cloudflareai" | "cloudflare" | "workersai" => { + Ok(Box::new(CloudflareAIAdapter::new(config)?)) + } // 阶段 2 适配器占位 - "azure" - | "anthropic" - | "runway" - | "pika" - | "kling" - | "stability" - | "chinese" - | "aggregation_platforms" - | "edge-tts" - | "elevenlabs" - | "cartesia" - | "deepgram" - | "assemblyai" => Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2 待实现)"), - }), + "anthropic" | "runway" | "pika" | "kling" | "stability" | "chinese" | "edge-tts" + | "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { + Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 2 待实现)"), + }) + } // 未知 provider _ => Err(AibridgeError::provider_not_found(format!( "{provider}(未知 provider,支持:{})", @@ -130,6 +138,78 @@ mod tests { assert_eq!(adapter.provider_type(), "gemini"); } + #[test] + fn create_azure_returns_adapter() { + // 阶段 2a:工厂已能构造真实 AzureAdapter(仅校验构造成功,不触发 HTTP) + // Azure 构造需 base_url 或 resource_name + deployment_id,此处用 base_url + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://example.openai.azure.com/openai/deployments/gpt-4") + .build(); + let config = ProviderConfig::from_options("azure", opts); + let adapter = create_adapter(config).expect("工厂应能创建 azure 适配器"); + assert_eq!(adapter.provider_type(), "azure"); + } + + #[test] + fn create_siliconflow_returns_adapter() { + let adapter = + create_adapter(config_for("siliconflow")).expect("工厂应能创建 siliconflow 适配器"); + assert_eq!(adapter.provider_type(), "siliconflow"); + } + + #[test] + fn create_togetherai_returns_adapter() { + let adapter = + create_adapter(config_for("togetherai")).expect("工厂应能创建 togetherai 适配器"); + assert_eq!(adapter.provider_type(), "togetherai"); + } + + #[test] + fn create_fireworksai_returns_adapter() { + let adapter = + create_adapter(config_for("fireworksai")).expect("工厂应能创建 fireworksai 适配器"); + assert_eq!(adapter.provider_type(), "fireworksai"); + } + + #[test] + fn create_cloudflareai_returns_adapter() { + // Cloudflare 需 extra.account_id 才能通过构造校验 + let opts = ClientOptions::builder() + .api_key("k") + .extra("account_id", "test-account") + .build(); + let config = ProviderConfig::from_options("cloudflareai", opts); + let adapter = create_adapter(config).expect("工厂应能创建 cloudflareai 适配器"); + assert_eq!(adapter.provider_type(), "cloudflareai"); + } + + #[test] + fn create_aggregation_platform_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/aggregation_platforms.py 末尾 register 调用: + // sf -> siliconflow / together -> togetherai / fireworks -> fireworksai + // cloudflare & workersai -> cloudflareai + let sf = create_adapter(config_for("sf")).expect("别名 sf 应映射到 siliconflow"); + assert_eq!(sf.provider_type(), "siliconflow"); + let together = + create_adapter(config_for("together")).expect("别名 together 应映射到 togetherai"); + assert_eq!(together.provider_type(), "togetherai"); + let fireworks = + create_adapter(config_for("fireworks")).expect("别名 fireworks 应映射到 fireworksai"); + assert_eq!(fireworks.provider_type(), "fireworksai"); + + let opts = ClientOptions::builder() + .api_key("k") + .extra("account_id", "test-account") + .build(); + let cloudflare = create_adapter(ProviderConfig::from_options("cloudflare", opts.clone())) + .expect("别名 cloudflare 应映射到 cloudflareai"); + assert_eq!(cloudflare.provider_type(), "cloudflareai"); + let workersai = create_adapter(ProviderConfig::from_options("workersai", opts)) + .expect("别名 workersai 应映射到 cloudflareai"); + assert_eq!(workersai.provider_type(), "cloudflareai"); + } + #[test] fn create_phase2_adapter_returns_phase2_message() { let result = create_adapter(config_for("anthropic")); @@ -146,6 +226,10 @@ mod tests { assert!(is_known_provider("openai")); assert!(is_known_provider("edge-tts")); assert!(is_known_provider("assemblyai")); + // 阶段 2a 已实现 provider 应被识别 + assert!(is_known_provider("azure")); + assert!(is_known_provider("siliconflow")); + assert!(is_known_provider("cloudflareai")); } #[test] diff --git a/crates/aibridge-core/src/adapters/gemini.rs b/crates/aibridge-core/src/adapters/gemini.rs index 55bf1af..489412a 100644 --- a/crates/aibridge-core/src/adapters/gemini.rs +++ b/crates/aibridge-core/src/adapters/gemini.rs @@ -2090,7 +2090,7 @@ mod tests { ); assert_eq!(result.get("topP").and_then(|v| v.as_f64()), Some(0.9)); assert_eq!(result.get("topK").and_then(|v| v.as_i64()), Some(40)); - assert!(result.get("stopSequences").is_some()); + assert!(result.contains_key("stopSequences")); // temperature 不在 rename_map,原名透传 assert_eq!( result.get("temperature").and_then(|v| v.as_f64()), diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 5128209..00f9914 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -27,3 +27,9 @@ pub mod openai; /// 火山引擎 CV 适配器:阶段 1 MVP,火山引擎视觉/视频生成协议 pub mod volcengine_cv; + +/// Azure OpenAI 适配器:阶段 2a,OpenAI 兼容协议的子适配器(Azure 部署) +pub mod azure; + +/// 聚合平台适配器:阶段 2a,含 SiliconFlow/TogetherAI/FireworksAI/CloudflareAI 四个 OpenAI 兼容子适配器 +pub mod aggregation_platforms; From 083a283783467811c26c15f0af1609a37d953242 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 20:50:52 +0800 Subject: [PATCH 25/55] =?UTF-8?q?chore:=20cargo=20fmt=20=E6=94=B6=E5=B0=BE?= =?UTF-8?q?=20aibridge-node/python=20lib.rs?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-node/src/lib.rs | 10 +-- crates/aibridge-python/src/lib.rs | 102 ++++++++++++++---------------- 2 files changed, 49 insertions(+), 63 deletions(-) diff --git a/crates/aibridge-node/src/lib.rs b/crates/aibridge-node/src/lib.rs index 37f924d..789c5ee 100644 --- a/crates/aibridge-node/src/lib.rs +++ b/crates/aibridge-node/src/lib.rs @@ -38,10 +38,7 @@ use aibridge_core::model::chat::{ChatCompletion, ChatCompletionChunk, ChatReques fn map_error(err: AibridgeError) -> Error { let code = err.code(); let message = err.to_string(); - Error::new( - Status::GenericFailure, - format!("[{code}] {message}"), - ) + Error::new(Status::GenericFailure, format!("[{code}] {message}")) } // ────────────────────────────────────────────────────────────────────────── @@ -361,10 +358,7 @@ impl Client { )) } }; - obj.insert( - "voice".into(), - serde_json::json!({ "voices": voices }), - ); + obj.insert("voice".into(), serde_json::json!({ "voices": voices })); } } diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs index 2ea809f..843e9d0 100644 --- a/crates/aibridge-python/src/lib.rs +++ b/crates/aibridge-python/src/lib.rs @@ -24,8 +24,8 @@ use std::sync::Arc; use futures::StreamExt; -use pyo3::prelude::*; use pyo3::coroutine::Coroutine; +use pyo3::prelude::*; use pyo3::types::{PyBytes, PyString}; use tokio::sync::Mutex; @@ -50,13 +50,12 @@ use aibridge_core::model::chat::{ /// 用 `once_cell::sync::Lazy` 在首次访问时初始化。core 的 async future(含 /// reqwest 等真实 IO)spawn 到此 runtime 上执行,PyO3 协程通过 await JoinHandle /// 取回结果。多线程 runtime 保证真实 adapter 的并发能力。 -static RUNTIME: once_cell::sync::Lazy = - once_cell::sync::Lazy::new(|| { - tokio::runtime::Builder::new_multi_thread() - .enable_all() - .build() - .expect("初始化 tokio runtime 失败") - }); +static RUNTIME: once_cell::sync::Lazy = once_cell::sync::Lazy::new(|| { + tokio::runtime::Builder::new_multi_thread() + .enable_all() + .build() + .expect("初始化 tokio runtime 失败") +}); // =========================================================================== // 错误映射 @@ -92,42 +91,17 @@ create_exception!( AibridgeError, "认证失败(API Key 无效/过期/无权限)" ); -create_exception!( - aibridge, - RateLimitError, - AibridgeError, - "请求频率超过限制" -); -create_exception!( - aibridge, - ValidationError, - AibridgeError, - "请求参数校验错误" -); +create_exception!(aibridge, RateLimitError, AibridgeError, "请求频率超过限制"); +create_exception!(aibridge, ValidationError, AibridgeError, "请求参数校验错误"); create_exception!( aibridge, ModelNotFoundError, AibridgeError, "请求的模型不存在" ); -create_exception!( - aibridge, - APIError, - AibridgeError, - "Provider API 调用错误" -); -create_exception!( - aibridge, - NetworkError, - AibridgeError, - "网络错误" -); -create_exception!( - aibridge, - TimeoutError, - AibridgeError, - "请求超时" -); +create_exception!(aibridge, APIError, AibridgeError, "Provider API 调用错误"); +create_exception!(aibridge, NetworkError, AibridgeError, "网络错误"); +create_exception!(aibridge, TimeoutError, AibridgeError, "请求超时"); create_exception!( aibridge, UnsupportedCapabilityError, @@ -213,7 +187,10 @@ impl ChatMessage { } fn __repr__(&self) -> String { - format!("ChatMessage(role={:?}, content={:?})", self.role, self.content) + format!( + "ChatMessage(role={:?}, content={:?})", + self.role, self.content + ) } } @@ -226,13 +203,11 @@ impl ChatMessage { if let Ok(msg) = obj.extract::>() { return Self::from_role_content(&msg.role, &msg.content); } - let dict: std::collections::HashMap = obj - .extract() - .map_err(|_| { - pyo3::exceptions::PyTypeError::new_err( - "消息必须是 ChatMessage 或含 role/content 的 dict", - ) - })?; + let dict: std::collections::HashMap = obj.extract().map_err(|_| { + pyo3::exceptions::PyTypeError::new_err( + "消息必须是 ChatMessage 或含 role/content 的 dict", + ) + })?; let role = dict .get("role") .ok_or_else(|| pyo3::exceptions::PyTypeError::new_err("消息缺少 role 字段"))?; @@ -360,7 +335,10 @@ impl ChoiceMessage { } fn __repr__(&self) -> String { - format!("ChoiceMessage(role={:?}, content={:?})", self.role, self.content) + format!( + "ChoiceMessage(role={:?}, content={:?})", + self.role, self.content + ) } } @@ -391,7 +369,10 @@ impl ChatCompletionChunk { } fn __repr__(&self) -> String { - format!("ChatCompletionChunk(id={:?}, model={:?})", self.id, self.model) + format!( + "ChatCompletionChunk(id={:?}, model={:?})", + self.id, self.model + ) } } @@ -447,7 +428,10 @@ impl ChatChunkDelta { } fn __repr__(&self) -> String { - format!("ChatChunkDelta(index={}, content={:?})", self.index, self.content) + format!( + "ChatChunkDelta(index={}, content={:?})", + self.index, self.content + ) } } @@ -502,7 +486,11 @@ impl SpeechResult { } fn __repr__(&self) -> String { - format!("SpeechResult(size={}, format={:?})", self.audio_data.len(), self.format) + format!( + "SpeechResult(size={}, format={:?})", + self.audio_data.len(), + self.format + ) } } @@ -795,9 +783,7 @@ impl Client { }) .await .map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!( - "chat_stream 任务失败: {e}" - )) + pyo3::exceptions::PyRuntimeError::new_err(format!("chat_stream 任务失败: {e}")) })?; let stream = stream_result.map_err(map_error)?; @@ -879,8 +865,14 @@ fn _aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { "UnsupportedCapabilityError", py.get_type::(), )?; - m.add("ProviderNotFoundError", py.get_type::())?; - m.add("VoiceNotAvailableError", py.get_type::())?; + m.add( + "ProviderNotFoundError", + py.get_type::(), + )?; + m.add( + "VoiceNotAvailableError", + py.get_type::(), + )?; m.add( "ServiceUnavailableError", py.get_type::(), From 65fab2f205f5fab61c1164071301dc84bae06454 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:04:50 +0800 Subject: [PATCH 26/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20additional=5Fmodels=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/adapters/additional_models.rs | 1801 +++++++++++++++++ 1 file changed, 1801 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/additional_models.rs diff --git a/crates/aibridge-core/src/adapters/additional_models.rs b/crates/aibridge-core/src/adapters/additional_models.rs new file mode 100644 index 0000000..9c0dfd1 --- /dev/null +++ b/crates/aibridge-core/src/adapters/additional_models.rs @@ -0,0 +1,1801 @@ +//! 更多主流模型适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/additional_models.py`。 +//! +//! 支持五个 OpenAI 兼容协议族的 Provider(差异主要在 base_url / provider_type / 能力声明): +//! - **xAI Grok**:`https://api.x.ai/v1`,chat + vision +//! - **零一万物 Yi**:`https://api.lingyiwanwu.com/v1`,chat + vision +//! - **商汤日日新 SenseNova**:`https://api.sensenova.cn/v1/cc-switch`,chat + vision, +//! list_models 端点为 `../llm/models`(base_url 末段是 `/cc-switch`,需相对路径上跳一级) +//! - **腾讯混元 Hunyuan**:`https://hunyuan.tencentcloudapi.com/v1`,chat + vision +//! - **Groq**:`https://api.groq.com/openai/v1`,chat + chat_stream + vision(Python 还声明 +//! AUDIO_TRANSCRIBE/TRANSLATE,但阶段 2a 范围仅核心对话能力,audio 走 trait 默认实现) +//! +//! ## 结构 +//! +//! Python v1 是 5 个独立 adapter 类,统一注册到工厂。Rust 同样实现 5 个独立 struct: +//! - `GrokAdapter` / `YiAdapter` / `HunyuanAdapter` / `GroqAdapter`:组合 `OpenAiCompatAdapter` +//! 地基,chat/chat_stream/list_models 全部委托(仅 base_url/provider_type/capabilities 差异) +//! - `SenseNovaAdapter`:chat/chat_stream 委托地基的请求体构造与响应解析,但 HTTP 层与 +//! list_models 端点独立实现(因 SenseNova 的 base_url 末段为 `/cc-switch`,list_models +//! 需用相对路径 `../llm/models` 上跳一级,与地基硬编码的 `models` 路径不同) +//! +//! ## provider_type 标识(对齐 Python v1) +//! +//! | Provider | provider_type | 别名 | +//! |---|---|---| +//! | xAI Grok | `grok` | `xaigrok` | +//! | 零一万物 Yi | `yi` | `lingyiwanwu` | +//! | 商汤日日新 SenseNova | `sensenova` | `shangtang` | +//! | 腾讯混元 Hunyuan | `hunyuan` | `tencent_hunyuan` | +//! | Groq | `groq` | — | +//! +//! ## 阶段范围 +//! +//! 阶段 2a 仅实现核心对话能力(chat / chat_stream)+ list_models 实时拉取。 +//! Python 老版中 Groq 声明的 audio transcribe/translate 与各 Provider 的 image/video +//! 不支持能力,均走 `Adapter` trait 默认实现返 `UnsupportedCapability`, +//! 待阶段 2c audio_adapters 统一补齐 audio 部分。 + +use async_trait::async_trait; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ChatCompletion, ChatRequest}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; + +// ==================== 默认 Base URL ==================== + +/// xAI Grok 默认 Base URL +/// +/// 对应 Python v1 `GrokAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_GROK_BASE_URL: &str = "https://api.x.ai/v1"; + +/// 零一万物 Yi 默认 Base URL +/// +/// 对应 Python v1 `YiAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_YI_BASE_URL: &str = "https://api.lingyiwanwu.com/v1"; + +/// 商汤日日新 SenseNova 默认 Base URL +/// +/// 对应 Python v1 `SenseNovaAdapter.DEFAULT_BASE_URL`。 +/// 末段 `/cc-switch` 是 chat 端点的 base,list_models 需上跳到 `/v1/llm/models`。 +pub const DEFAULT_SENSENOVA_BASE_URL: &str = "https://api.sensenova.cn/v1/cc-switch"; + +/// 腾讯混元 Hunyuan 默认 Base URL +/// +/// 对应 Python v1 `HunyuanAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_HUNYUAN_BASE_URL: &str = "https://hunyuan.tencentcloudapi.com/v1"; + +/// Groq 默认 Base URL +/// +/// 对应 Python v1 `GroqAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_GROQ_BASE_URL: &str = "https://api.groq.com/openai/v1"; + +// ==================== 能力集合构造 ==================== + +/// xAI Grok 支持的能力集合 +/// +/// 对齐 Python v1 `GrokAdapter.supported_capabilities = ["chat", "vision"]`。 +/// chat_stream 虽未在 Python 显式声明,但 Python 实现了该方法且 OpenAI 兼容协议天然支持流式, +/// 故 Rust 一并声明 ChatStream(与 openai/azure 等兼容适配器保持一致)。 +fn grok_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 零一万物 Yi 支持的能力集合 +/// +/// 对齐 Python v1 `YiAdapter.supported_capabilities = ["chat", "vision"]`。 +fn yi_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 商汤日日新 SenseNova 支持的能力集合 +/// +/// 对齐 Python v1 `SenseNovaAdapter.supported_capabilities = ["chat", "vision"]`。 +fn sensenova_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 腾讯混元 Hunyuan 支持的能力集合 +/// +/// 对齐 Python v1 `HunyuanAdapter.supported_capabilities = ["chat", "vision"]`。 +fn hunyuan_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// Groq 支持的能力集合 +/// +/// 对齐 Python v1 `GroqAdapter.supported_capabilities`(CHAT / CHAT_STREAM / VISION + +/// AUDIO_TRANSCRIBE / AUDIO_TRANSLATE)。阶段 2a 范围仅核心对话能力, +/// audio 走 trait 默认实现返 UnsupportedCapability,待阶段 2c 补齐。 +fn groq_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +// ==================== xAI Grok 适配器 ==================== + +/// xAI Grok 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.x.ai/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct GrokAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl GrokAdapter { + /// 创建 xAI Grok 适配器 + /// + /// `config.base_url` 为空时回退到 [`DEFAULT_GROK_BASE_URL`]。 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "grok", + "xAI Grok", + DEFAULT_GROK_BASE_URL, + grok_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 地基构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for GrokAdapter { + fn provider_type(&self) -> &str { + "grok" + } + + fn provider_name(&self) -> &str { + "xAI Grok" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / embed / video / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 零一万物 Yi 适配器 ==================== + +/// 零一万物 Yi 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.lingyiwanwu.com/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct YiAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl YiAdapter { + /// 创建零一万物 Yi 适配器 + /// + /// `config.base_url` 为空时回退到 [`DEFAULT_YI_BASE_URL`]。 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "yi", + "零一万物 Yi", + DEFAULT_YI_BASE_URL, + yi_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 地基构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for YiAdapter { + fn provider_type(&self) -> &str { + "yi" + } + + fn provider_name(&self) -> &str { + "零一万物 Yi" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== 商汤日日新 SenseNova 适配器 ==================== + +/// 商汤日日新 SenseNova 适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream 复用地基的请求体构造与响应解析, +/// 但 list_models 端点独立实现: +/// - Base URL: `https://api.sensenova.cn/v1/cc-switch`(末段 `/cc-switch` 是 chat 端点 base) +/// - Chat: `POST /chat/completions`(即 `.../v1/cc-switch/chat/completions`) +/// - Models: `GET ../llm/models`(相对路径上跳一级到 `.../v1/llm/models`) +/// - 认证: Bearer Token +/// +/// 因地基地 `list_models` 硬编码相对路径 `models`(会拼成 `.../cc-switch/models`,错误), +/// 故本适配器持有一个独立 HttpClient,自行实现 list_models; +/// chat/chat_stream 仍委托地基(地基的 `chat/completions` 相对路径与 SenseNova 一致)。 +pub struct SenseNovaAdapter { + /// OpenAI 兼容地基(chat/chat_stream 委托给它,复用请求体构造与响应解析) + compat: OpenAiCompatAdapter, + /// list_models 端点专用 HTTP 客户端(独立于 compat,避免暴露 compat 私有字段) + http: HttpClient, + /// list_models 端点的完整 URL(已上跳一级,形如 `.../v1/llm/models`) + /// + /// 由 base_url 末段去掉 `/cc-switch` 后拼 `llm/models` 得到。 + /// 用 String 而非每次计算,避免重复字符串处理。 + models_url: String, + /// API key(Bearer 认证用,可能为空) + api_key: Option, +} + +impl SenseNovaAdapter { + /// 创建商汤日日新 SenseNova 适配器 + /// + /// 解析顺序(与 Python v1 `__init__` 一致): + /// 1. `config.base_url` 为空时用 [`DEFAULT_SENSENOVA_BASE_URL`] 兜底 + /// 2. 由 base_url 计算 models_url:取 base_url 去掉末段(`/cc-switch`)后拼 `llm/models` + /// + /// # models_url 计算示例 + /// - base_url = `https://api.sensenova.cn/v1/cc-switch` + /// - models_url = `https://api.sensenova.cn/v1/llm/models` + /// + /// # 错误 + /// - HTTP 客户端构造失败(罕见,reqwest 配置错误) + pub fn new(config: ProviderConfig) -> Result { + let caps = sensenova_capabilities(); + + // 解析 base_url:config 优先,否则默认值 + let base_url = config + .base_url + .clone() + .map(|s| s.trim().to_string()) + .filter(|s| !s.is_empty()) + .unwrap_or_else(|| DEFAULT_SENSENOVA_BASE_URL.to_string()); + + // 由 base_url 计算 models_url:去掉末段路径,拼 `llm/models` + // base_url 形如 `.../v1/cc-switch`,去掉最后一段得到 `.../v1`,再拼 `llm/models` + let models_url = build_sensenova_models_url(&base_url); + + // 构造 chat 用的 compat 地基(base_url 保持原值,chat/completions 相对路径匹配) + let compat = OpenAiCompatAdapter::new( + config.clone(), + "sensenova", + "商汤日日新 SenseNova", + &base_url, + caps, + )?; + + // 构造 list_models 用的 HttpClient(不设 base_url,请求时用完整 models_url) + let http_opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&http_opts)?; + + Ok(Self { + compat, + http, + models_url, + api_key: config.api_key.clone(), + }) + } + + /// 用显式 compat + http 构造(测试用,可注入 mockito 后端) + /// + /// `models_url` 为 list_models 的完整 URL;`api_key` 用于 Bearer 认证。 + #[cfg(test)] + pub fn with_compat( + compat: OpenAiCompatAdapter, + http: HttpClient, + models_url: String, + api_key: Option, + ) -> Self { + Self { + compat, + http, + models_url, + api_key, + } + } + + /// list_models 端点完整 URL(已上跳一级,形如 `.../v1/llm/models`) + pub fn models_url(&self) -> &str { + &self.models_url + } + + /// 发送带 Bearer 认证的 GET 请求,并用 OpenAI 错误映射处理响应 + /// + /// SenseNova 错误体结构与 OpenAI 一致(`{"error": {"message": "..."}}`), + /// 故复用 [`OpenAiCompatAdapter::map_api_error`]。 + async fn get_authed_json(&self, url: &str) -> Result { + let resp = self + .http + .inner() + .get(url) + .bearer_auth(self.api_key.as_deref().unwrap_or("")) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::() + .await + .map_err(AibridgeError::from) + } +} + +/// 由 SenseNova chat base_url 计算 list_models 端点 URL +/// +/// base_url 末段为 `/cc-switch`(chat 端点 base),需上跳一级拼 `llm/models`。 +/// 例:`https://api.sensenova.cn/v1/cc-switch` → `https://api.sensenova.cn/v1/llm/models`。 +/// +/// 实现方式:找到最后一个 `/`,截到其前一位;若末段已是空(以 `/` 结尾),则直接拼 `llm/models`。 +/// 保留 scheme + host 段不变(只处理 path 部分)。 +fn build_sensenova_models_url(base_url: &str) -> String { + // 去掉末尾斜杠,统一处理 + let trimmed = base_url.trim_end_matches('/'); + // 截到倒数第二个 `/`(即去掉最后一段 path) + // 例:`https://host/v1/cc-switch` → `https://host/v1` + let parent = match trimmed.rfind('/') { + Some(idx) => &trimmed[..idx], + None => trimmed, + }; + format!("{parent}/llm/models") +} + +#[async_trait] +impl Adapter for SenseNovaAdapter { + fn provider_type(&self) -> &str { + "sensenova" + } + + fn provider_name(&self) -> &str { + "商汤日日新 SenseNova" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + /// + /// SenseNova 的 base_url 末段 `/cc-switch` + 地基相对路径 `chat/completions` + /// 正好拼出 `.../v1/cc-switch/chat/completions`,与官方端点一致。 + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取) + /// + /// 调用 `GET {models_url}`(即 `.../v1/llm/models`),解析响应。 + /// 响应结构与 OpenAI `/models` 一致:`{"data": [{"id": "...", ...}]}`。 + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json(&self.models_url).await?; + let models = parse_models_response(&value, "sensenova"); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 腾讯混元 Hunyuan 适配器 ==================== + +/// 腾讯混元 Hunyuan 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://hunyuan.tencentcloudapi.com/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +pub struct HunyuanAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl HunyuanAdapter { + /// 创建腾讯混元 Hunyuan 适配器 + /// + /// `config.base_url` 为空时回退到 [`DEFAULT_HUNYUAN_BASE_URL`]。 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "hunyuan", + "腾讯混元 Hunyuan", + DEFAULT_HUNYUAN_BASE_URL, + hunyuan_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 地基构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for HunyuanAdapter { + fn provider_type(&self) -> &str { + "hunyuan" + } + + fn provider_name(&self) -> &str { + "腾讯混元 Hunyuan" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Groq 适配器 ==================== + +/// Groq 适配器 +/// +/// OpenAI 兼容协议,全部核心能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.groq.com/openai/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +/// +/// Python v1 还声明了 AUDIO_TRANSCRIBE / AUDIO_TRANSLATE(Whisper 语音识别), +/// 阶段 2a 范围仅核心对话能力,audio 走 trait 默认实现返 UnsupportedCapability, +/// 待阶段 2c audio_adapters 补齐。 +pub struct GroqAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl GroqAdapter { + /// 创建 Groq 适配器 + /// + /// `config.base_url` 为空时回退到 [`DEFAULT_GROQ_BASE_URL`]。 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "groq", + "Groq", + DEFAULT_GROQ_BASE_URL, + groq_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 地基构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for GroqAdapter { + fn provider_type(&self) -> &str { + "groq" + } + + fn provider_name(&self) -> &str { + "Groq" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致(speech 显式抛错也走默认实现)。 +} + +// ==================== 内部:模型列表响应解析 ==================== + +/// 解析 OpenAI 兼容 `/models`(或 SenseNova `/llm/models`)响应 → Vec +/// +/// 响应格式:`{"data": [{"id": "...", "created": ..., "owned_by": "..."}]}` +/// 与 [`OpenAiCompatAdapter`] 内部的 `parse_models` 等价,但因该方法是私有,本模块独立实现一份。 +/// +/// `provider` 填入对应 provider_type(如 "sensenova"),用于 ModelInfo.provider 字段。 +fn parse_models_response(value: &serde_json::Value, provider: &str) -> Vec { + let arr = value.get("data").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .map(|m| { + let id = m + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let model_type = infer_model_type(&id); + ModelInfo { + name: id.clone(), + id, + model_type, + provider: provider.to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: None, + created: m.get("created").and_then(|v| v.as_u64()), + } + }) + .collect(), + None => Vec::new(), + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::error::AibridgeError; + use crate::http::HttpClient; + use crate::model::chat::ChatMessage; + use crate::model::image::ImageRequest; + use crate::model::video::VideoRequest; + use futures::stream::StreamExt; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 OpenAiCompatAdapter(指向 mockito server,给定 provider 信息与能力) + fn make_compat( + server: &Server, + provider_type: &str, + provider_name: &str, + caps: CapabilitySet, + ) -> OpenAiCompatAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options(provider_type, opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + OpenAiCompatAdapter::with_http(http, config, provider_type, provider_name, caps) + } + + /// 构造测试用 GrokAdapter(指向 mockito server) + fn make_grok(server: &Server) -> GrokAdapter { + let compat = make_compat(server, "grok", "xAI Grok", grok_capabilities()); + GrokAdapter::with_compat(compat) + } + + /// 构造测试用 YiAdapter(指向 mockito server) + fn make_yi(server: &Server) -> YiAdapter { + let compat = make_compat(server, "yi", "零一万物 Yi", yi_capabilities()); + YiAdapter::with_compat(compat) + } + + /// 构造测试用 SenseNovaAdapter(指向 mockito server) + /// + /// mockito server 的 URL 作为 chat base_url(形如 `http://127.0.0.1:port`), + /// chat 端点走 compat 地基,请求 `{server}/chat/completions`; + /// list_models 走独立 http,请求 `{server}/llm/models`(直接覆盖 models_url 以隔离计算逻辑)。 + fn make_sensenova(server: &Server) -> SenseNovaAdapter { + let compat = make_compat( + server, + "sensenova", + "商汤日日新 SenseNova", + sensenova_capabilities(), + ); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let models_url = format!("{}/llm/models", server.url()); + SenseNovaAdapter::with_compat(compat, http, models_url, Some("test-key".to_string())) + } + + /// 构造测试用 HunyuanAdapter(指向 mockito server) + fn make_hunyuan(server: &Server) -> HunyuanAdapter { + let compat = make_compat( + server, + "hunyuan", + "腾讯混元 Hunyuan", + hunyuan_capabilities(), + ); + HunyuanAdapter::with_compat(compat) + } + + /// 构造测试用 GroqAdapter(指向 mockito server) + fn make_groq(server: &Server) -> GroqAdapter { + let compat = make_compat(server, "groq", "Groq", groq_capabilities()); + GroqAdapter::with_compat(compat) + } + + /// 构造不指向任何 server 的 GrokAdapter(用于不发请求的元信息/能力测试) + fn make_grok_no_server() -> GrokAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_GROK_BASE_URL) + .build(); + let config = ProviderConfig::from_options("grok", opts); + GrokAdapter::new(config).expect("GrokAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 YiAdapter + fn make_yi_no_server() -> YiAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_YI_BASE_URL) + .build(); + let config = ProviderConfig::from_options("yi", opts); + YiAdapter::new(config).expect("YiAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 HunyuanAdapter + fn make_hunyuan_no_server() -> HunyuanAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_HUNYUAN_BASE_URL) + .build(); + let config = ProviderConfig::from_options("hunyuan", opts); + HunyuanAdapter::new(config).expect("HunyuanAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 GroqAdapter + fn make_groq_no_server() -> GroqAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_GROQ_BASE_URL) + .build(); + let config = ProviderConfig::from_options("groq", opts); + GroqAdapter::new(config).expect("GroqAdapter 构造应成功") + } + + // ============ Grok 元信息 ============ + + #[test] + fn grok_provider_type_and_name_match_python() { + let adapter = make_grok_no_server(); + assert_eq!(adapter.provider_type(), "grok"); + assert_eq!(adapter.provider_name(), "xAI Grok"); + } + + #[test] + fn grok_requires_api_key_is_true() { + let adapter = make_grok_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn grok_capabilities_contains_chat_and_vision() { + let adapter = make_grok_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // image / video / audio 不声明(阶段 2a 范围) + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::VideoGenerate)); + } + + #[test] + fn grok_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("grok", opts); + let adapter = GrokAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_GROK_BASE_URL); + } + + #[test] + fn grok_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.grok-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("grok", opts); + let adapter = GrokAdapter::new(config).unwrap(); + assert_eq!( + adapter.compat.base_url(), + "https://custom.grok-proxy.com/v1" + ); + } + + // ============ Grok chat ============ + + #[tokio::test] + async fn grok_chat_success_returns_completion() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "grok-3", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.model, "grok-3"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn grok_chat_sends_temperature_and_max_tokens() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "grok-3", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "grok-3", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn grok_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid xAI API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn grok_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn grok_chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ Grok chat_stream ============ + + #[tokio::test] + async fn grok_chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"grok-3\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"grok-3\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"grok-3\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn grok_chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let req = ChatRequest::builder("grok-3", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + // ============ Grok list_models ============ + + #[tokio::test] + async fn grok_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(json!({ + "object": "list", + "data": [ + {"id": "grok-3", "object": "model", "created": 1700000000, "owned_by": "xai"}, + {"id": "grok-2-vision", "object": "model", "created": 1700000000, "owned_by": "xai"} + ] + }).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "grok-3"); + assert_eq!(models[0].provider, "grok"); + assert_eq!(models[1].id, "grok-2-vision"); + } + + #[tokio::test] + async fn grok_list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_grok(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ Grok 不支持的能力 ============ + + #[tokio::test] + async fn grok_image_generate_returns_unsupported() { + let adapter = make_grok_no_server(); + let req = ImageRequest::builder("grok-2-image", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn grok_video_create_returns_unsupported() { + let adapter = make_grok_no_server(); + let req = VideoRequest::builder("grok-video", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ Yi 元信息 + chat + list_models ============ + + #[test] + fn yi_provider_type_and_name_match_python() { + let adapter = make_yi_no_server(); + assert_eq!(adapter.provider_type(), "yi"); + assert_eq!(adapter.provider_name(), "零一万物 Yi"); + } + + #[test] + fn yi_requires_api_key_is_true() { + let adapter = make_yi_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn yi_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("yi", opts); + let adapter = YiAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_YI_BASE_URL); + } + + #[tokio::test] + async fn yi_chat_success_returns_completion() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "id": "chatcmpl-yi", + "object": "chat.completion", + "created": 1700000000, + "model": "yi-large", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "你好!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 3, "completion_tokens": 3, "total_tokens": 6} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_yi(&server); + let req = ChatRequest::builder("yi-large", vec![ChatMessage::user("你好")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.model, "yi-large"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("你好!")); + } + + #[tokio::test] + async fn yi_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid Yi API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_yi(&server); + let req = ChatRequest::builder("yi-large", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn yi_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "yi-large", "object": "model", "created": 1, "owned_by": "01ai"}, + {"id": "yi-vision", "object": "model", "created": 1, "owned_by": "01ai"} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_yi(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "yi-large"); + assert_eq!(models[0].provider, "yi"); + } + + #[tokio::test] + async fn yi_list_models_filter_by_chat_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "yi-large", "object": "model", "created": 1}, + {"id": "yi-vision", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_yi(&server); + // yi-large 和 yi-vision 都推断为 Chat 类型(infer_model_type 按名称关键字推断) + let chats = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert_eq!(chats.len(), 2); + } + + // ============ SenseNova 元信息 + chat + list_models ============ + + #[test] + fn sensenova_provider_type_and_name_match_python() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url(DEFAULT_SENSENOVA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("sensenova", opts); + let adapter = SenseNovaAdapter::new(config).expect("SenseNovaAdapter 构造应成功"); + assert_eq!(adapter.provider_type(), "sensenova"); + assert_eq!(adapter.provider_name(), "商汤日日新 SenseNova"); + } + + #[test] + fn sensenova_requires_api_key_is_true() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url(DEFAULT_SENSENOVA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("sensenova", opts); + let adapter = SenseNovaAdapter::new(config).unwrap(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn sensenova_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("sensenova", opts); + let adapter = SenseNovaAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_SENSENOVA_BASE_URL); + } + + #[test] + fn sensenova_models_url_builds_from_default_base() { + // 默认 base_url = https://api.sensenova.cn/v1/cc-switch + // models_url 应为 https://api.sensenova.cn/v1/llm/models + let opts = ClientOptions::builder() + .api_key("k") + .base_url(DEFAULT_SENSENOVA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("sensenova", opts); + let adapter = SenseNovaAdapter::new(config).unwrap(); + assert_eq!( + adapter.models_url(), + "https://api.sensenova.cn/v1/llm/models" + ); + } + + #[test] + fn sensenova_models_url_builds_from_custom_base_with_trailing_slash() { + // 自定义 base_url(含末尾斜杠也应正确处理) + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.sensenova.cn/v2/cc-switch/") + .build(); + let config = ProviderConfig::from_options("sensenova", opts); + let adapter = SenseNovaAdapter::new(config).unwrap(); + assert_eq!( + adapter.models_url(), + "https://custom.sensenova.cn/v2/llm/models" + ); + } + + #[tokio::test] + async fn sensenova_chat_success_returns_completion() { + let mut server = Server::new_async().await; + // chat 端点走 compat 地基,相对路径 chat/completions 拼到 server 根 + server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "id": "chatcmpl-sn", + "object": "chat.completion", + "created": 1700000000, + "model": "SenseChat-5", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "你好"}, + "finish_reason": "stop" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let req = ChatRequest::builder("SenseChat-5", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.model, "SenseChat-5"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("你好")); + } + + #[tokio::test] + async fn sensenova_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let req = ChatRequest::builder("SenseChat-5", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn sensenova_chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"SenseChat-5\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"你好\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"SenseChat-5\",\"choices\":[{\"index\":0,\"delta\":{},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let req = ChatRequest::builder("SenseChat-5", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 2); + assert_eq!(chunks[0].choices[0].delta.content.as_deref(), Some("你好")); + } + + #[tokio::test] + async fn sensenova_list_models_success() { + let mut server = Server::new_async().await; + // list_models 走独立 http,请求 {server}/llm/models + server + .mock("GET", "/llm/models") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "SenseChat-5", "object": "model", "created": 1}, + {"id": "SenseChat-Vision", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "SenseChat-5"); + assert_eq!(models[0].provider, "sensenova"); + } + + #[tokio::test] + async fn sensenova_list_models_filter_by_image_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/llm/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "SenseChat-5", "object": "model", "created": 1}, + {"id": "dall-e-3", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn sensenova_list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/llm/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn sensenova_list_models_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("GET", "/llm/models") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_sensenova(&server); + let err = adapter.list_models(None).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ Hunyuan 元信息 + chat + list_models ============ + + #[test] + fn hunyuan_provider_type_and_name_match_python() { + let adapter = make_hunyuan_no_server(); + assert_eq!(adapter.provider_type(), "hunyuan"); + assert_eq!(adapter.provider_name(), "腾讯混元 Hunyuan"); + } + + #[test] + fn hunyuan_requires_api_key_is_true() { + let adapter = make_hunyuan_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn hunyuan_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("hunyuan", opts); + let adapter = HunyuanAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_HUNYUAN_BASE_URL); + } + + #[tokio::test] + async fn hunyuan_chat_success_returns_completion() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "id": "chatcmpl-hy", + "object": "chat.completion", + "created": 1700000000, + "model": "hunyuan-pro", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "你好!"}, + "finish_reason": "stop" + }] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_hunyuan(&server); + let req = ChatRequest::builder("hunyuan-pro", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.model, "hunyuan-pro"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("你好!")); + } + + #[tokio::test] + async fn hunyuan_chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body(json!({"error": {"message": "model hunyuan-x not found"}}).to_string()) + .create_async() + .await; + + let adapter = make_hunyuan(&server); + let req = ChatRequest::builder("hunyuan-x", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn hunyuan_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "hunyuan-pro", "object": "model", "created": 1}, + {"id": "hunyuan-standard", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_hunyuan(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "hunyuan-pro"); + assert_eq!(models[0].provider, "hunyuan"); + } + + #[tokio::test] + async fn hunyuan_list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow"}}).to_string()) + .create_async() + .await; + + let adapter = make_hunyuan(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ============ Groq 元信息 + chat + list_models ============ + + #[test] + fn groq_provider_type_and_name_match_python() { + let adapter = make_groq_no_server(); + assert_eq!(adapter.provider_type(), "groq"); + assert_eq!(adapter.provider_name(), "Groq"); + } + + #[test] + fn groq_requires_api_key_is_true() { + let adapter = make_groq_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn groq_capabilities_contains_chat_stream() { + let adapter = make_groq_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // audio 阶段 2a 不声明(走默认实现) + assert!(!caps.contains(&Capabilities::AudioTranscribe)); + assert!(!caps.contains(&Capabilities::AudioSpeech)); + } + + #[test] + fn groq_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("groq", opts); + let adapter = GroqAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_GROQ_BASE_URL); + } + + #[tokio::test] + async fn groq_chat_success_returns_completion() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "id": "chatcmpl-groq", + "object": "chat.completion", + "created": 1700000000, + "model": "llama-3.1-70b-versatile", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hi!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 2, "completion_tokens": 2, "total_tokens": 4} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_groq(&server); + let req = + ChatRequest::builder("llama-3.1-70b-versatile", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.model, "llama-3.1-70b-versatile"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hi!")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 4); + } + + #[tokio::test] + async fn groq_chat_passes_extra_params_through() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "llama-3.1-70b-versatile", + "custom_param": "custom_value" + }))) + .with_status(200) + .with_body(json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "llama-3.1-70b-versatile", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }).to_string()) + .create_async() + .await; + + let adapter = make_groq(&server); + let req = ChatRequest::builder("llama-3.1-70b-versatile", vec![ChatMessage::user("hi")]) + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn groq_chat_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(400) + .with_body(json!({"error": {"message": "max_tokens is invalid"}}).to_string()) + .create_async() + .await; + + let adapter = make_groq(&server); + let req = + ChatRequest::builder("llama-3.1-70b-versatile", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("max_tokens")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn groq_chat_stream_sends_stream_true() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "stream": true, + "stream_options": {"include_usage": true} + }))) + .with_status(200) + .with_body("data: [DONE]\n") + .create_async() + .await; + + let adapter = make_groq(&server); + let req = + ChatRequest::builder("llama-3.1-70b-versatile", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + #[tokio::test] + async fn groq_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body(json!({ + "data": [ + {"id": "llama-3.1-70b-versatile", "object": "model", "created": 1, "owned_by": "groq"}, + {"id": "whisper-large-v3", "object": "model", "created": 1, "owned_by": "groq"} + ] + }).to_string()) + .create_async() + .await; + + let adapter = make_groq(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "llama-3.1-70b-versatile"); + assert_eq!(models[0].provider, "groq"); + // whisper 推断为 Audio 类型 + assert_eq!(models[1].model_type, ModelType::Audio); + } + + #[tokio::test] + async fn groq_list_models_filter_audio() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "llama-3.1-70b-versatile", "object": "model", "created": 1}, + {"id": "whisper-large-v3", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_groq(&server); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 1); + assert_eq!(audio[0].id, "whisper-large-v3"); + } + + // ============ Groq 不支持的能力(阶段 2a 范围)============ + + #[tokio::test] + async fn groq_speech_returns_unsupported() { + // Python v1 Groq.speech 显式抛 UnsupportedCapabilityError(仅支持 Whisper ASR,无 TTS) + let adapter = make_groq_no_server(); + let req = crate::model::audio::SpeechRequest::builder("tts-1", "hi", "alloy").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn groq_image_generate_returns_unsupported() { + let adapter = make_groq_no_server(); + let req = ImageRequest::builder("groq-image", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ start / close ============ + + #[tokio::test] + async fn grok_start_and_close_are_noops() { + let mut adapter = make_grok_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn sensenova_start_and_close_are_noops() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url(DEFAULT_SENSENOVA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("sensenova", opts); + let mut adapter = SenseNovaAdapter::new(config).unwrap(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn groq_start_and_close_are_noops() { + let mut adapter = make_groq_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ build_sensenova_models_url 单元测试 ============ + + #[test] + fn build_sensenova_models_url_default() { + let url = build_sensenova_models_url("https://api.sensenova.cn/v1/cc-switch"); + assert_eq!(url, "https://api.sensenova.cn/v1/llm/models"); + } + + #[test] + fn build_sensenova_models_url_trailing_slash() { + let url = build_sensenova_models_url("https://api.sensenova.cn/v1/cc-switch/"); + assert_eq!(url, "https://api.sensenova.cn/v1/llm/models"); + } + + #[test] + fn build_sensenova_models_url_custom_version() { + let url = build_sensenova_models_url("https://custom.sensenova.cn/v2/cc-switch"); + assert_eq!(url, "https://custom.sensenova.cn/v2/llm/models"); + } +} From 3fd6df3421e477d9bea3c1960c5ef5107ded70fd Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:07:37 +0800 Subject: [PATCH 27/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20more=5Fmodels=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../aibridge-core/src/adapters/more_models.rs | 2175 +++++++++++++++++ .../src/adapters/openai_compat.rs | 24 + 2 files changed, 2199 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/more_models.rs diff --git a/crates/aibridge-core/src/adapters/more_models.rs b/crates/aibridge-core/src/adapters/more_models.rs new file mode 100644 index 0000000..0d98735 --- /dev/null +++ b/crates/aibridge-core/src/adapters/more_models.rs @@ -0,0 +1,2175 @@ +//! 更多主流模型适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/more_models.py`。 +//! +//! 支持五个主流模型 Provider(文档见各适配器注释): +//! - **DeepSeek**:OpenAI 兼容协议,支持思考模式(reasoning_effort + 自动注入 thinking) +//! - **阶跃星辰 StepFun**:OpenAI 兼容协议(base_url 含 `/v1` 前缀) +//! - **Mistral AI**:OpenAI 兼容协议(base_url 含 `/v1` 前缀) +//! - **Cohere**:非标准协议(`POST /chat`,message/chat_history 结构,响应 `text` 字段) +//! - **Perplexity AI**:OpenAI 兼容协议,AI 搜索 +//! +//! ## 结构 +//! +//! Python v1 是 5 个独立 adapter 类,统一注册到工厂(StepFun 额外注册 `step` 别名)。 +//! Rust 同样实现 5 个独立 struct: +//! - `DeepSeekAdapter` / `StepFunAdapter` / `MistralAdapter` / `PerplexityAdapter`: +//! 组合 `OpenAiCompatAdapter` 地基,chat/chat_stream/image/embed/list_models 全部委托 +//! (DeepSeek 额外在 chat 前注入 `thinking`,对齐 Python 自动思考模式行为) +//! - `CohereAdapter`:组合 `OpenAiCompatAdapter` 仅复用 embed 能力校验等公共逻辑, +//! chat/chat_stream/list_models 独立实现(Cohere 特有 `/chat` 端点 + message/chat_history +//! 请求体 + `text` 响应 + `event_type` 流式事件 + `/models` 响应 `models[].name` 字段) +//! +//! ## provider_type 标识(对齐 Python v1) +//! +//! | Provider | provider_type | 别名 | +//! |---|---|---| +//! | DeepSeek | `deepseek` | — | +//! | 阶跃星辰 StepFun | `stepfun` | `step` | +//! | Mistral AI | `mistral` | — | +//! | Cohere | `cohere` | — | +//! | Perplexity AI | `perplexity` | — | +//! +//! 注册到工厂由 `adapter::factory` 完成(阶段 2a 收尾批次),本模块仅声明适配器类型。 + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChoiceMessage, DeltaMessage, UserContent, +}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::util; + +// ==================== 默认 Base URL ==================== + +/// DeepSeek 默认 Base URL +/// +/// 对应 Python v1 `DeepSeekAdapter.DEFAULT_BASE_URL`。 +/// DeepSeek base_url 不含 `/v1` 前缀,chat 端点为 `POST /chat/completions`, +/// models 端点为 `GET /models`(与 OpenAI 兼容地基的相对路径一致)。 +pub const DEFAULT_DEEPSEEK_BASE_URL: &str = "https://api.deepseek.com"; + +/// 阶跃星辰 StepFun 默认 Base URL +/// +/// 对应 Python v1 `StepFunAdapter.DEFAULT_BASE_URL = "https://api.stepfun.com"`, +/// 但 Python `start()` 中 httpx base_url 设为 `base_url + "/v1"`,故实际请求 base +/// 为 `https://api.stepfun.com/v1`。此处直接采用含 `/v1` 的值,行为等价。 +pub const DEFAULT_STEPFUN_BASE_URL: &str = "https://api.stepfun.com/v1"; + +/// Mistral AI 默认 Base URL +/// +/// 对应 Python v1 `MistralAdapter.DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +pub const DEFAULT_MISTRAL_BASE_URL: &str = "https://api.mistral.ai/v1"; + +/// Cohere 默认 Base URL +/// +/// 对应 Python v1 `CohereAdapter.DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +/// Cohere 的 chat 端点为 `POST /chat`(非标准 OpenAI),models 端点为 `GET /models`。 +pub const DEFAULT_COHERE_BASE_URL: &str = "https://api.cohere.ai/v1"; + +/// Perplexity AI 默认 Base URL +/// +/// 对应 Python v1 `PerplexityAdapter.DEFAULT_BASE_URL`。 +/// Perplexity base_url 不含 `/v1` 前缀,chat 端点为 `POST /chat/completions`, +/// models 端点为 `GET /models`(与 OpenAI 兼容地基的相对路径一致)。 +pub const DEFAULT_PERPLEXITY_BASE_URL: &str = "https://api.perplexity.ai"; + +// ==================== 能力集合构造 ==================== + +/// DeepSeek 支持的能力集合 +/// +/// 对齐 Python v1 `DeepSeekAdapter.supported_capabilities = ["chat", "vision"]`。 +/// Rust 额外声明 ChatStream(OpenAI 兼容地基支持流式,Python 老版也实现了 chat_stream)。 +fn deepseek_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// StepFun 支持的能力集合 +/// +/// 对齐 Python v1 `StepFunAdapter.supported_capabilities = ["chat", "vision"]`。 +fn stepfun_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// Mistral 支持的能力集合 +/// +/// 对齐 Python v1 `MistralAdapter.supported_capabilities = ["chat", "vision"]`。 +fn mistral_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// Cohere 支持的能力集合 +/// +/// 对齐 Python v1 `CohereAdapter.supported_capabilities = ["chat", "vision"]`。 +fn cohere_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// Perplexity 支持的能力集合 +/// +/// 对齐 Python v1 `PerplexityAdapter.supported_capabilities = ["chat", "vision"]`。 +/// Perplexity 的 Sonar 模型具备联网搜索能力,额外声明 WebSearch。 +fn perplexity_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps.insert(Capabilities::WebSearch); + caps +} + +// ==================== DeepSeek 适配器 ==================== + +/// DeepSeek 适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream/image(不支持)/embed(不支持)/list_models 委托给 +/// `OpenAiCompatAdapter` 地基。DeepSeek 特有行为:当请求体含 `reasoning_effort` 时 +/// 自动注入 `thinking: {"type": "enabled"}`(对齐 Python v1 `_build_request_body`)。 +/// +/// - Base URL: `https://api.deepseek.com` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +/// - 文档: +pub struct DeepSeekAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl DeepSeekAdapter { + /// 创建 DeepSeek 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_DEEPSEEK_BASE_URL, + deepseek_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "deepseek"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "DeepSeek"; + + /// 构造 DeepSeek chat 请求体(含自动 thinking 注入) + /// + /// 对齐 Python v1 `DeepSeekAdapter._build_request_body`: + /// 1. 先用兼容地基构造标准 OpenAI 请求体(model/messages/temperature/max_tokens/top_p 等) + /// 2. 若 body 含 `reasoning_effort` 字段(来自 ChatRequest.reasoning_effort 或 extra 透传), + /// 且 body 尚无 `thinking` 字段,则自动注入 `thinking: {"type": "enabled"}` + fn build_chat_body(&self, req: &ChatRequest, stream: bool) -> Value { + let mut body = self.compat.build_chat_body(req, stream); + // 自动思考模式:有 reasoning_effort 且无 thinking 时注入 + if body.get("reasoning_effort").is_some() && body.get("thinking").is_none() { + if let Some(obj) = body.as_object_mut() { + obj.insert("thinking".to_string(), json!({ "type": "enabled" })); + } + } + body + } +} + +#[async_trait] +impl Adapter for DeepSeekAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话:构造 DeepSeek 请求体(含自动 thinking),委托地基发送 + 解析 + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.ensure_capability_pub(Capabilities::Chat)?; + let body = self.build_chat_body(&req, false); + let value = self.compat.post_chat(&body).await?; + self.compat.parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability,与 Python 老版行为一致。 +} + +// ==================== StepFun 适配器 ==================== + +/// 阶跃星辰 StepFun 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.stepfun.com/v1`(Python default 不含 `/v1`,但 `start()` 拼接, +/// 此处直接采用含 `/v1` 的值,行为等价) +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +/// - 文档: +pub struct StepFunAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl StepFunAdapter { + /// 创建 StepFun 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_STEPFUN_BASE_URL, + stepfun_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "stepfun"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "阶跃星辰 StepFun"; +} + +#[async_trait] +impl Adapter for StepFunAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Mistral 适配器 ==================== + +/// Mistral AI 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// +/// - Base URL: `https://api.mistral.ai/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +/// - 文档: +pub struct MistralAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl MistralAdapter { + /// 创建 Mistral 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_MISTRAL_BASE_URL, + mistral_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "mistral"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Mistral AI"; +} + +#[async_trait] +impl Adapter for MistralAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +// ==================== Cohere 适配器 ==================== + +/// Cohere 适配器 +/// +/// 非标准 OpenAI 协议,chat/chat_stream/list_models 独立实现,不委托兼容地基。 +/// +/// - Base URL: `https://api.cohere.ai/v1` +/// - Chat: `POST /chat`(请求体含 `message` + `chat_history` + `system_prompt`,非标准 OpenAI) +/// - 响应: `{"text": "...", "usage": {"tokens": {"input_tokens":..., "output_tokens":...}}}` +/// - 流式: SSE,每行 `data: {"event_type": "text-generation"|"stream-end", "text": "..."}` +/// - Models: `GET /models`(响应 `{"models": [{"name": "...", ...}]}`,`name` 即模型 ID) +/// - 认证: Bearer Token +/// - 文档: +/// +/// embed 走 trait 默认实现返 UnsupportedCapability(Python v1 未实现 embed)。 +pub struct CohereAdapter { + /// OpenAI 兼容地基(仅复用其 HttpClient 与错误映射,不委托 chat 等方法) + compat: OpenAiCompatAdapter, +} + +impl CohereAdapter { + /// 创建 Cohere 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_COHERE_BASE_URL, + cohere_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "cohere"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Cohere"; + + /// API key + fn api_key(&self) -> Option<&str> { + self.compat.api_key() + } + + /// base_url + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), Self::PROVIDER_TYPE), + }) + } + } + + /// 发送带认证的 POST JSON 请求,并用 OpenAI 错误映射处理响应 + /// + /// 复用兼容地基的 `map_api_error`,错误分类与 OpenAI 兼容族一致。 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .compat + .http_inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .header("accept", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带认证的 GET 请求,并用 OpenAI 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .compat + .http_inner() + .get(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 转换统一消息为 Cohere 格式,并提取 system prompt + /// + /// 对齐 Python v1 `CohereAdapter._convert_messages`: + /// - system 消息 → 提取为 `system_prompt`(最后一条 system 覆盖前面的) + /// - user 消息 → `{"role": "USER", "content": ...}`(多模态取文本拼接,纯文本直接取) + /// - 其他(assistant 等)→ `{"role": "CHATBOT", "content": ...}` + /// + /// 返回 `(chat_history, system_prompt)`,`chat_history` 不含最后一条 user 消息 + /// (最后一条 user 消息作为 Cohere `message` 字段单独发送)。 + fn convert_messages(messages: &[ChatMessage]) -> (Vec, Option) { + let mut converted: Vec = Vec::new(); + let mut system_prompt: Option = None; + + for msg in messages { + match msg { + ChatMessage::System { content, .. } => { + system_prompt = Some(content.clone()); + } + ChatMessage::User { content, .. } => { + let text = match content { + UserContent::Text(s) => s.clone(), + // 多模态:拼接所有 Text 部件,忽略图片 + UserContent::Parts(parts) => parts + .iter() + .filter_map(|p| match p { + crate::model::chat::ContentPart::Text { text } => { + Some(text.clone()) + } + _ => None, + }) + .collect::>() + .join(""), + }; + converted.push(json!({ "role": "USER", "content": text })); + } + ChatMessage::Assistant { content, .. } => { + let text = content.clone().unwrap_or_default(); + converted.push(json!({ "role": "CHATBOT", "content": text })); + } + ChatMessage::Tool { content, .. } => { + // 工具结果消息按 CHATBOT 角色处理(Cohere 无原生 tool 角色) + converted.push(json!({ "role": "CHATBOT", "content": content })); + } + } + } + + (converted, system_prompt) + } + + /// 构造 Cohere chat 请求体 + /// + /// 对齐 Python v1 `CohereAdapter.chat` 的 body 构造: + /// - `message`:最后一条 user 消息的 content(无则空串) + /// - `chat_history`:除最后一条 user 消息外的所有消息 + /// - `system_prompt`:若有 system 消息则填入 + /// - `temperature` / `max_tokens`:从统一请求透传 + /// - `extra` 透传到顶层 + fn build_chat_body(req: &ChatRequest, stream: bool) -> Value { + let (converted, system_prompt) = Self::convert_messages(&req.messages); + + // 最后一条 user 消息的 content 作为 Cohere message 字段 + let message = converted + .last() + .and_then(|m| m.get("content")) + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + // chat_history 不含最后一条消息 + let chat_history = if converted.len() > 1 { + converted[..converted.len() - 1].to_vec() + } else { + Vec::new() + }; + + let mut body = json!({ + "model": req.model, + "message": message, + "chat_history": chat_history, + }); + if stream { + body["stream"] = json!(true); + } + if let Some(sp) = system_prompt { + body["system_prompt"] = json!(sp); + } + if let Some(t) = req.temperature { + body["temperature"] = json!(t); + } + if let Some(mt) = req.max_tokens { + body["max_tokens"] = json!(mt); + } + // extra 透传到顶层 + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 解析 Cohere chat 响应 → ChatCompletion + /// + /// Cohere 响应格式:`{"text": "...", "usage": {"tokens": {"input_tokens":..., "output_tokens":...}}}` + /// (非标准 OpenAI choices 结构)。对应 Python v1 `CohereAdapter._parse_response`。 + fn parse_chat_completion(value: &Value, fallback_model: &str) -> Result { + let text = value + .get("text") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + + // usage 解析:Cohere 用 tokens.input_tokens / tokens.output_tokens + let usage = value.get("usage").and_then(parse_cohere_usage); + + Ok(ChatCompletion { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion".to_string(), + created: value + .get("created_at") + .and_then(|v| v.as_u64()) + .or_else(|| value.get("created").and_then(|v| v.as_u64())) + .unwrap_or_else(util::current_timestamp), + model: fallback_model.to_string(), + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".to_string(), + content: Some(text), + tool_calls: None, + }, + finish_reason: Some("stop".to_string()), + }], + usage, + service_tier: None, + system_fingerprint: None, + }) + } + + /// 解析单个 Cohere 流式 chunk + /// + /// Cohere 流式事件格式:`{"event_type": "text-generation"|"stream-end", "text": "...", ...}` + /// - `text-generation`:增量文本,delta.content = text,finish_reason = None + /// - `stream-end`:结束事件,delta.content = "",finish_reason = "stop" + /// - 其他 event_type:返回 None(跳过) + /// + /// 对应 Python v1 `CohereAdapter._parse_chunk`。 + fn parse_chunk(value: &Value, fallback_model: &str) -> Option { + let event_type = value + .get("event_type") + .and_then(|v| v.as_str()) + .unwrap_or(""); + let generation_id = value + .get("generation_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")); + + match event_type { + "text-generation" => { + let text = value.get("text").and_then(|v| v.as_str()).unwrap_or(""); + Some(ChatCompletionChunk { + id: generation_id, + object: "chat.completion.chunk".to_string(), + created: util::current_timestamp(), + model: fallback_model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".to_string()), + content: Some(text.to_string()), + tool_calls: None, + }, + finish_reason: None, + }], + usage: None, + }) + } + "stream-end" => Some(ChatCompletionChunk { + id: generation_id, + object: "chat.completion.chunk".to_string(), + created: util::current_timestamp(), + model: fallback_model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".to_string()), + content: Some(String::new()), + tool_calls: None, + }, + finish_reason: Some("stop".to_string()), + }], + usage: None, + }), + _ => None, + } + } + + /// 解析 Cohere /models 响应 → Vec + /// + /// Cohere 响应:`{"models": [{"name": "...", ...}]}`,模型 ID 字段名为 `name` + /// (非标准 OpenAI `id`),需转换为统一 `id` 字段。 + /// 对应 Python v1 `CohereAdapter.list_models` 的预处理逻辑。 + fn parse_models(value: &Value, provider: &str) -> Vec { + let arr = value.get("models").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .map(|m| { + // Cohere 用 name 作为模型 ID,统一映射到 id + let id = m + .get("id") + .and_then(|v| v.as_str()) + .or_else(|| m.get("name").and_then(|v| v.as_str())) + .unwrap_or("") + .to_string(); + let model_type = infer_model_type(&id); + ModelInfo { + name: id.clone(), + id, + model_type, + provider: provider.to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: m + .get("description") + .and_then(|v| v.as_str()) + .map(str::to_owned), + created: None, + } + }) + .collect(), + None => Vec::new(), + } + } +} + +#[async_trait] +impl Adapter for CohereAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话(Cohere 特有协议:POST /chat) + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = Self::build_chat_body(&req, false); + let value = self.post_authed_json("chat", &body).await?; + Self::parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话(Cohere 特有协议:POST /chat stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = Self::build_chat_body(&req, true); + let url = self.url("chat"); + + let resp = self + .compat + .http_inner() + .post(&url) + .bearer_auth(self.api_key().unwrap_or("")) + .header("accept", "application/json") + .json(&body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(OpenAiCompatAdapter::map_api_error(status_code, &body_text)); + } + + let model = req.model.clone(); + // 按字节流读取,按行切分解析 SSE(Cohere 流式同样是 SSE data: 格式) + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = CohereLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 去除 "data: " 前缀 + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + // 结束标记 + if data.trim() == "[DONE]" { + return; + } + match serde_json::from_str::(data) { + Ok(v) => match Self::parse_chunk(&v, &model) { + Some(chunk) => yield Ok(chunk), + None => continue, + }, + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 模型列表(Cohere 特有:GET /models,响应 models[].name 即模型 ID) + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("models").await?; + let models = Self::parse_models(&value, self.provider_type()); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability,与 Python 老版行为一致。 +} + +/// 解析 Cohere usage 统计 +/// +/// Cohere usage 格式:`{"tokens": {"input_tokens": N, "output_tokens": M}}`, +/// 需转换为统一 ChatUsage(prompt/completion/total)。 +fn parse_cohere_usage(v: &Value) -> Option { + let tokens = v.get("tokens")?; + let prompt = tokens.get("input_tokens").and_then(|x| x.as_u64())?; + let completion = tokens + .get("output_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(0); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: prompt + completion, + }) +} + +// ==================== Cohere SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(Cohere 流式用) +/// +/// 与 `openai_compat::LinesStream` 等价,独立实现避免引用其私有结构。 +struct CohereLinesStream { + inner: S, + buffer: Vec, +} + +impl CohereLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for CohereLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + // 先看缓冲区是否已有完整行 + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + // 去掉末尾 \n 与可能的 \r + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + // 缓冲区无完整行,拉取下一 chunk + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + // 继续循环,尝试从缓冲区切出行 + } + Poll::Ready(None) => { + // 流结束,把缓冲区剩余内容作为最后一行返回 + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +// ==================== Perplexity 适配器 ==================== + +/// Perplexity AI 适配器 +/// +/// OpenAI 兼容协议,全部能力委托给 `OpenAiCompatAdapter` 地基。 +/// Perplexity 特有参数 `extra_body` 经统一请求的 `extra` 透传到请求体顶层。 +/// +/// - Base URL: `https://api.perplexity.ai` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models` +/// - 认证: Bearer Token +/// - 文档: +pub struct PerplexityAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl PerplexityAdapter { + /// 创建 Perplexity 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_PERPLEXITY_BASE_URL, + perplexity_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "perplexity"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Perplexity AI"; +} + +#[async_trait] +impl Adapter for PerplexityAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::http::HttpClient; + use crate::model::chat::ChatMessage; + use crate::model::image::ImageRequest; + use crate::model::options::{EmbedInput, EmbedRequest, ReasoningEffort}; + use crate::model::video::VideoRequest; + use futures::stream::StreamExt; + use mockito::Server; + use std::collections::HashMap; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 OpenAiCompatAdapter(指向 mockito server,给定 provider 信息与能力) + fn make_compat( + server: &Server, + provider_type: &str, + provider_name: &str, + caps: CapabilitySet, + ) -> OpenAiCompatAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options(provider_type, opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + OpenAiCompatAdapter::with_http(http, config, provider_type, provider_name, caps) + } + + /// 构造测试用 DeepSeekAdapter(指向 mockito server) + fn make_deepseek(server: &Server) -> DeepSeekAdapter { + let compat = make_compat(server, "deepseek", "DeepSeek", deepseek_capabilities()); + DeepSeekAdapter::with_compat(compat) + } + + /// 构造测试用 StepFunAdapter(指向 mockito server) + fn make_stepfun(server: &Server) -> StepFunAdapter { + let compat = make_compat( + server, + "stepfun", + "阶跃星辰 StepFun", + stepfun_capabilities(), + ); + StepFunAdapter::with_compat(compat) + } + + /// 构造测试用 MistralAdapter(指向 mockito server) + fn make_mistral(server: &Server) -> MistralAdapter { + let compat = make_compat(server, "mistral", "Mistral AI", mistral_capabilities()); + MistralAdapter::with_compat(compat) + } + + /// 构造测试用 CohereAdapter(指向 mockito server) + fn make_cohere(server: &Server) -> CohereAdapter { + let compat = make_compat(server, "cohere", "Cohere", cohere_capabilities()); + CohereAdapter::with_compat(compat) + } + + /// 构造测试用 PerplexityAdapter(指向 mockito server) + fn make_perplexity(server: &Server) -> PerplexityAdapter { + let compat = make_compat( + server, + "perplexity", + "Perplexity AI", + perplexity_capabilities(), + ); + PerplexityAdapter::with_compat(compat) + } + + /// 标准 OpenAI chat 成功响应体 + fn openai_chat_body() -> Value { + json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "test-model", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }) + } + + // ==================== 默认 Base URL 常量校验 ==================== + + #[test] + fn default_base_urls_match_python() { + assert_eq!(DEFAULT_DEEPSEEK_BASE_URL, "https://api.deepseek.com"); + assert_eq!(DEFAULT_STEPFUN_BASE_URL, "https://api.stepfun.com/v1"); + assert_eq!(DEFAULT_MISTRAL_BASE_URL, "https://api.mistral.ai/v1"); + assert_eq!(DEFAULT_COHERE_BASE_URL, "https://api.cohere.ai/v1"); + assert_eq!(DEFAULT_PERPLEXITY_BASE_URL, "https://api.perplexity.ai"); + } + + // ==================== DeepSeek 测试 ==================== + + #[tokio::test] + async fn deepseek_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + assert_eq!(adapter.provider_type(), "deepseek"); + assert_eq!(adapter.provider_name(), "DeepSeek"); + } + + #[tokio::test] + async fn deepseek_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn deepseek_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // 不支持 image / embed + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::Embedding)); + } + + #[tokio::test] + async fn deepseek_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("deepseek-chat", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn deepseek_chat_auto_injects_thinking_when_reasoning_effort_present() { + // 对齐 Python:传 reasoning_effort 时自动注入 thinking={"type":"enabled"} + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "deepseek-reasoner", + "reasoning_effort": "high", + "thinking": {"type": "enabled"} + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("deepseek-reasoner", vec![ChatMessage::user("hi")]) + .reasoning_effort(ReasoningEffort::High) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn deepseek_chat_no_thinking_without_reasoning_effort() { + // 不传 reasoning_effort 时不应注入 thinking。 + // 用 PartialJson 验证请求体不含 thinking 字段:mockito 无原生"不含字段"断言, + // 故用完整 JsonString 严格匹配预期 body(model + messages + temperature,无 thinking)。 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "deepseek-chat", + "messages": [{"role": "user", "content": "hi"}] + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("deepseek-chat", vec![ChatMessage::user("hi")]).build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn deepseek_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "invalid key"}}).to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn deepseek_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn deepseek_chat_stream_parses_sse() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"deepseek-chat\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"deepseek-chat\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"deepseek-chat\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + let adapter = make_deepseek(&server); + let req = ChatRequest::builder("deepseek-chat", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + } + + #[tokio::test] + async fn deepseek_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "deepseek-chat", "object": "model"}, + {"id": "deepseek-reasoner", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_deepseek(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "deepseek-chat"); + assert_eq!(models[0].provider, "deepseek"); + } + + #[tokio::test] + async fn deepseek_list_models_filter_by_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "deepseek-chat", "object": "model"}, + {"id": "dall-e-3", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_deepseek(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn deepseek_list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_deepseek(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn deepseek_image_generate_unsupported() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + let req = ImageRequest::builder("m", "prompt").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn deepseek_video_create_unsupported() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + let req = VideoRequest::builder("m", "p").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn deepseek_embed_unsupported() { + let server = Server::new_async().await; + let adapter = make_deepseek(&server); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== StepFun 测试 ==================== + + #[tokio::test] + async fn stepfun_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_stepfun(&server); + assert_eq!(adapter.provider_type(), "stepfun"); + assert_eq!(adapter.provider_name(), "阶跃星辰 StepFun"); + } + + #[tokio::test] + async fn stepfun_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "step-1-8k", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_stepfun(&server); + let req = ChatRequest::builder("step-1-8k", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn stepfun_chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body(json!({"error": {"message": "model not found"}}).to_string()) + .create_async() + .await; + let adapter = make_stepfun(&server); + let req = ChatRequest::builder("bad-model", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn stepfun_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "step-1-8k", "object": "model"}, + {"id": "step-1v-8k", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_stepfun(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].provider, "stepfun"); + } + + #[tokio::test] + async fn stepfun_chat_stream_sends_stream_true() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "stream": true, + "stream_options": {"include_usage": true} + }))) + .with_status(200) + .with_body("data: [DONE]\n") + .create_async() + .await; + let adapter = make_stepfun(&server); + let req = ChatRequest::builder("step-1-8k", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + // ==================== Mistral 测试 ==================== + + #[tokio::test] + async fn mistral_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_mistral(&server); + assert_eq!(adapter.provider_type(), "mistral"); + assert_eq!(adapter.provider_name(), "Mistral AI"); + } + + #[tokio::test] + async fn mistral_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "mistral-large-latest", + "temperature": 0.7, + "top_p": 0.9 + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_mistral(&server); + let req = ChatRequest::builder("mistral-large-latest", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .top_p(0.9) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn mistral_chat_error_429() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "rate limit"}}).to_string()) + .create_async() + .await; + let adapter = make_mistral(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn mistral_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "mistral-large-latest", "object": "model"}, + {"id": "mistral-embed", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_mistral(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].provider, "mistral"); + } + + #[tokio::test] + async fn mistral_chat_stream_parses_sse() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"mistral-large-latest\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hi\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"mistral-large-latest\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"!\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_mistral(&server); + let req = + ChatRequest::builder("mistral-large-latest", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut content = String::new(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.unwrap(); + if let Some(c) = chunk.choices[0].delta.content.as_deref() { + content.push_str(c); + } + } + assert_eq!(content, "Hi!"); + } + + // ==================== Cohere 测试 ==================== + + #[tokio::test] + async fn cohere_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_cohere(&server); + assert_eq!(adapter.provider_type(), "cohere"); + assert_eq!(adapter.provider_name(), "Cohere"); + } + + #[tokio::test] + async fn cohere_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_cohere(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn cohere_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat") + .match_header("authorization", "Bearer test-key") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "command-r-plus", + "message": "hi", + "chat_history": [], + "temperature": 0.7 + }))) + .with_status(200) + .with_body( + json!({ + "id": "gen-1", + "text": "Hello from Cohere!", + "usage": {"tokens": {"input_tokens": 3, "output_tokens": 4}} + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("command-r-plus", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "gen-1"); + assert_eq!(resp.choices.len(), 1); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("Hello from Cohere!") + ); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + // usage 解析:Cohere tokens.input/output → prompt/completion/total + let usage = resp.usage.as_ref().expect("应有 usage"); + assert_eq!(usage.prompt_tokens, 3); + assert_eq!(usage.completion_tokens, 4); + assert_eq!(usage.total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn cohere_chat_extracts_system_prompt_and_history() { + // 多消息场景:system 提取为 system_prompt,历史消息进 chat_history, + // 最后一条 user 作为 message + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "command-r-plus", + "message": "second question", + "chat_history": [ + {"role": "USER", "content": "first question"}, + {"role": "CHATBOT", "content": "first answer"} + ], + "system_prompt": "you are helpful" + }))) + .with_status(200) + .with_body(json!({"text": "ok"}).to_string()) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder( + "command-r-plus", + vec![ + ChatMessage::system("you are helpful"), + ChatMessage::user("first question"), + ChatMessage::assistant("first answer"), + ChatMessage::user("second question"), + ], + ) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("ok")); + mock.assert_async().await; + } + + #[tokio::test] + async fn cohere_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat") + .with_status(401) + .with_body(json!({"error": {"message": "invalid token"}}).to_string()) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn cohere_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn cohere_chat_stream_parses_events() { + let mut server = Server::new_async().await; + // Cohere 流式:text-generation 事件 + stream-end 事件 + [DONE] + let sse = "data: {\"event_type\":\"text-generation\",\"text\":\"Hello\",\"generation_id\":\"g1\"}\n\ + data: {\"event_type\":\"text-generation\",\"text\":\" world\",\"generation_id\":\"g1\"}\n\ + data: {\"event_type\":\"stream-end\",\"generation_id\":\"g1\"}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("command-r-plus", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + // 3 个有效 chunk:2 个 text-generation + 1 个 stream-end + assert_eq!(chunks.len(), 3); + // 拼接前两块内容 + let mut content = String::new(); + content.push_str(chunks[0].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + // 第三块是 stream-end,finish_reason = stop,content 为空 + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(chunks[2].choices[0].delta.content.as_deref(), Some("")); + } + + #[tokio::test] + async fn cohere_chat_stream_skips_unknown_event_types() { + let mut server = Server::new_async().await; + // 含未知 event_type(stream-start)应被跳过 + let sse = "data: {\"event_type\":\"stream-start\"}\n\ + data: {\"event_type\":\"text-generation\",\"text\":\"Hi\"}\n\ + data: {\"event_type\":\"stream-end\"}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat") + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + // stream-start 被跳过,仅 text-generation + stream-end + assert_eq!(chunks.len(), 2); + } + + #[tokio::test] + async fn cohere_chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat") + .with_status(401) + .with_body(json!({"error": {"message": "bad token"}}).to_string()) + .create_async() + .await; + let adapter = make_cohere(&server); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn cohere_list_models_success_converts_name_to_id() { + // Cohere /models 响应用 name 字段作为模型 ID,需转换为统一 id + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "models": [ + {"name": "command-r-plus", "description": "Command R+"}, + {"name": "command-r", "description": "Command R"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cohere(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + // name 字段被映射到 id + assert_eq!(models[0].id, "command-r-plus"); + assert_eq!(models[0].name, "command-r-plus"); + assert_eq!(models[0].provider, "cohere"); + assert_eq!(models[0].description.as_deref(), Some("Command R+")); + } + + #[tokio::test] + async fn cohere_list_models_filter_by_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "models": [ + {"name": "command-r-plus"}, + {"name": "dall-e-3"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_cohere(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "dall-e-3"); + } + + #[tokio::test] + async fn cohere_list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad token"}}).to_string()) + .create_async() + .await; + let adapter = make_cohere(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn cohere_image_generate_unsupported() { + let server = Server::new_async().await; + let adapter = make_cohere(&server); + let req = ImageRequest::builder("m", "prompt").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn cohere_embed_unsupported() { + let server = Server::new_async().await; + let adapter = make_cohere(&server); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== Perplexity 测试 ==================== + + #[tokio::test] + async fn perplexity_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_perplexity(&server); + assert_eq!(adapter.provider_type(), "perplexity"); + assert_eq!(adapter.provider_name(), "Perplexity AI"); + } + + #[tokio::test] + async fn perplexity_capabilities_include_web_search() { + let server = Server::new_async().await; + let adapter = make_perplexity(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::WebSearch)); + } + + #[tokio::test] + async fn perplexity_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "sonar", + "temperature": 0.5 + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_perplexity(&server); + let req = ChatRequest::builder("sonar", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn perplexity_chat_passes_extra_body_through() { + // Perplexity 特有参数 extra_body 经统一请求的 extra 透传到请求体顶层 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "sonar", + "extra_body": {"search_recency_filter": "month"} + }))) + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_perplexity(&server); + let req = ChatRequest::builder("sonar", vec![ChatMessage::user("hi")]) + .extra("extra_body", json!({"search_recency_filter": "month"})) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn perplexity_chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(404) + .with_body(json!({"error": {"message": "model not found"}}).to_string()) + .create_async() + .await; + let adapter = make_perplexity(&server); + let req = ChatRequest::builder("bad-model", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn perplexity_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "sonar", "object": "model"}, + {"id": "sonar-pro", "object": "model"} + ] + }) + .to_string(), + ) + .create_async() + .await; + let adapter = make_perplexity(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "sonar"); + assert_eq!(models[0].provider, "perplexity"); + } + + #[tokio::test] + async fn perplexity_chat_stream_parses_sse() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"sonar\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Answer\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"sonar\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"!\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_perplexity(&server); + let req = ChatRequest::builder("sonar", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut content = String::new(); + while let Some(chunk) = stream.next().await { + let chunk = chunk.unwrap(); + if let Some(c) = chunk.choices[0].delta.content.as_deref() { + content.push_str(c); + } + } + assert_eq!(content, "Answer!"); + } + + #[tokio::test] + async fn perplexity_image_generate_unsupported() { + let server = Server::new_async().await; + let adapter = make_perplexity(&server); + let req = ImageRequest::builder("m", "prompt").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== start/close 无副作用 ==================== + + #[tokio::test] + async fn deepseek_start_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_deepseek(&server); + adapter.start().await.unwrap(); + adapter.close().await.unwrap(); + } + + #[tokio::test] + async fn cohere_start_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_cohere(&server); + adapter.start().await.unwrap(); + adapter.close().await.unwrap(); + } + + // ==================== Cohere 辅助函数单元测试 ==================== + + #[test] + fn cohere_parse_chunk_text_generation_returns_content() { + let v = json!({"event_type": "text-generation", "text": "hello"}); + let chunk = CohereAdapter::parse_chunk(&v, "command-r").expect("应返回 Some"); + assert_eq!(chunk.choices[0].delta.content.as_deref(), Some("hello")); + assert!(chunk.choices[0].finish_reason.is_none()); + } + + #[test] + fn cohere_parse_chunk_stream_end_returns_stop() { + let v = json!({"event_type": "stream-end", "generation_id": "g1"}); + let chunk = CohereAdapter::parse_chunk(&v, "command-r").expect("应返回 Some"); + assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(chunk.choices[0].delta.content.as_deref(), Some("")); + assert_eq!(chunk.id, "g1"); + } + + #[test] + fn cohere_parse_chunk_unknown_event_returns_none() { + let v = json!({"event_type": "stream-start"}); + assert!(CohereAdapter::parse_chunk(&v, "command-r").is_none()); + } + + #[test] + fn cohere_parse_models_handles_name_field() { + let v = json!({ + "models": [ + {"name": "command-r-plus", "description": "top"}, + {"id": "embed-english-v3.0"} + ] + }); + let models = CohereAdapter::parse_models(&v, "cohere"); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "command-r-plus"); + assert_eq!(models[0].description.as_deref(), Some("top")); + // 第二个用 id 字段 + assert_eq!(models[1].id, "embed-english-v3.0"); + } + + #[test] + fn cohere_parse_models_empty_returns_vec() { + let v = json!({"models": []}); + let models = CohereAdapter::parse_models(&v, "cohere"); + assert!(models.is_empty()); + // 无 models 字段 + let v2 = json!({}); + assert!(CohereAdapter::parse_models(&v2, "cohere").is_empty()); + } + + #[test] + fn cohere_parse_chat_completion_extracts_text_and_usage() { + let v = json!({ + "id": "gen-9", + "text": "response text", + "usage": {"tokens": {"input_tokens": 10, "output_tokens": 20}} + }); + let cc = CohereAdapter::parse_chat_completion(&v, "command-r").unwrap(); + assert_eq!(cc.id, "gen-9"); + assert_eq!( + cc.choices[0].message.content.as_deref(), + Some("response text") + ); + let usage = cc.usage.unwrap(); + assert_eq!(usage.prompt_tokens, 10); + assert_eq!(usage.completion_tokens, 20); + assert_eq!(usage.total_tokens, 30); + } + + #[test] + fn cohere_parse_chat_completion_without_usage() { + let v = json!({"text": "no usage here"}); + let cc = CohereAdapter::parse_chat_completion(&v, "command-r").unwrap(); + assert_eq!( + cc.choices[0].message.content.as_deref(), + Some("no usage here") + ); + assert!(cc.usage.is_none()); + } + + #[test] + fn cohere_convert_messages_extracts_system_and_history() { + let msgs = vec![ + ChatMessage::system("be helpful"), + ChatMessage::user("q1"), + ChatMessage::assistant("a1"), + ChatMessage::user("q2"), + ]; + let (history, system) = CohereAdapter::convert_messages(&msgs); + assert_eq!(system.as_deref(), Some("be helpful")); + // chat_history 应含前 3 条(system 已提取,user q1 + assistant a1) + assert_eq!(history.len(), 3); + assert_eq!(history[0]["role"], "USER"); + assert_eq!(history[0]["content"], "q1"); + assert_eq!(history[1]["role"], "CHATBOT"); + assert_eq!(history[1]["content"], "a1"); + // 最后一条 user 作为 message,在 build_chat_body 中单独处理 + assert_eq!(history[2]["role"], "USER"); + assert_eq!(history[2]["content"], "q2"); + } +} diff --git a/crates/aibridge-core/src/adapters/openai_compat.rs b/crates/aibridge-core/src/adapters/openai_compat.rs index 6552ac7..a4a3c5f 100644 --- a/crates/aibridge-core/src/adapters/openai_compat.rs +++ b/crates/aibridge-core/src/adapters/openai_compat.rs @@ -165,6 +165,30 @@ impl OpenAiCompatAdapter { &self.capabilities } + /// 内部 reqwest::Client 引用(供子适配器独立实现非标准端点时复用 HTTP 客户端) + /// + /// 对子适配器开放(pub):Cohere 等非标准协议适配器需直接用 reqwest::Client + /// 构造非标准请求(如 POST /chat + 自定义 body),复用本结构的连接池与超时配置。 + pub fn http_inner(&self) -> &reqwest::Client { + self.http.inner() + } + + /// 校验请求的能力是否被支持(对子适配器开放的 pub 版本) + /// + /// 私有 `ensure_capability` 的 pub 包装,供 DeepSeek/Cohere 等子适配器在 + /// 重写 chat 等方法时复用能力校验逻辑。 + pub fn ensure_capability_pub(&self, cap: Capabilities) -> Result<()> { + self.ensure_capability(cap) + } + + /// 发送 chat/completions 请求(对子适配器开放) + /// + /// 私有 `post_authed_json("chat/completions", body)` 的 pub 包装, + /// 供 DeepSeek 等子适配器在自定义请求体后复用发送 + 错误映射逻辑。 + pub async fn post_chat(&self, body: &Value) -> Result { + self.post_authed_json("chat/completions", body).await + } + /// 拼接完整 URL(base_url + 相对路径) fn url(&self, path: &str) -> String { let base = self.base_url.trim_end_matches('/'); From a0eb96dee38d9fb231a53b475ca8700ee967d148 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:14:01 +0800 Subject: [PATCH 28/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20=E7=AC=AC=E4=BA=8C=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20additional=5Fmodels=20+=20more=5Fmodels=20=E5=88=B0?= =?UTF-8?q?=E5=B7=A5=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 127 ++++++++++++++++++++ crates/aibridge-core/src/adapters/mod.rs | 6 + 2 files changed, 133 insertions(+) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index 8b0c89d..a87efec 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -8,6 +8,9 @@ //! - 阶段 0.4 暂只占位分支(返 ProviderNotFound),具体适配器阶段 1 起填充 use crate::adapter::Adapter; +use crate::adapters::additional_models::{ + GrokAdapter, GroqAdapter, HunyuanAdapter, SenseNovaAdapter, YiAdapter, +}; use crate::adapters::aggregation_platforms::{ CloudflareAIAdapter, FireworksAIAdapter, SiliconFlowAdapter, TogetherAIAdapter, }; @@ -15,6 +18,9 @@ use crate::adapters::agnes::AgnesAdapter; use crate::adapters::azure::AzureAdapter; use crate::adapters::echo::EchoAdapter; use crate::adapters::gemini::GeminiAdapter; +use crate::adapters::more_models::{ + CohereAdapter, DeepSeekAdapter, MistralAdapter, PerplexityAdapter, StepFunAdapter, +}; use crate::adapters::openai::OpenAiAdapter; use crate::adapters::volcengine_cv::VolcengineCvAdapter; use crate::config::ProviderConfig; @@ -36,6 +42,16 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "togetherai", "fireworksai", "cloudflareai", + "grok", + "yi", + "sensenova", + "hunyuan", + "groq", + "deepseek", + "stepfun", + "mistral", + "cohere", + "perplexity", // 阶段 2b/2c 待实现: "anthropic", "runway", @@ -74,6 +90,18 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { "cloudflareai" | "cloudflare" | "workersai" => { Ok(Box::new(CloudflareAIAdapter::new(config)?)) } + // 扩展模型:别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用 + "grok" | "xaigrok" => Ok(Box::new(GrokAdapter::new(config)?)), + "yi" | "lingyiwanwu" => Ok(Box::new(YiAdapter::new(config)?)), + "sensenova" | "shangtang" => Ok(Box::new(SenseNovaAdapter::new(config)?)), + "hunyuan" | "tencent_hunyuan" => Ok(Box::new(HunyuanAdapter::new(config)?)), + "groq" => Ok(Box::new(GroqAdapter::new(config)?)), + // 更多模型:别名对齐 Python agn/adapters/more_models.py 末尾 register 调用 + "deepseek" => Ok(Box::new(DeepSeekAdapter::new(config)?)), + "stepfun" | "step" => Ok(Box::new(StepFunAdapter::new(config)?)), + "mistral" => Ok(Box::new(MistralAdapter::new(config)?)), + "cohere" => Ok(Box::new(CohereAdapter::new(config)?)), + "perplexity" => Ok(Box::new(PerplexityAdapter::new(config)?)), // 阶段 2 适配器占位 "anthropic" | "runway" | "pika" | "kling" | "stability" | "chinese" | "edge-tts" | "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { @@ -210,6 +238,94 @@ mod tests { assert_eq!(workersai.provider_type(), "cloudflareai"); } + #[test] + fn create_grok_returns_adapter() { + // 阶段 2a additional_models:GrokAdapter 自带 DEFAULT_GROK_BASE_URL 回退,仅需 api_key + let adapter = create_adapter(config_for("grok")).expect("工厂应能创建 grok 适配器"); + assert_eq!(adapter.provider_type(), "grok"); + } + + #[test] + fn create_yi_returns_adapter() { + let adapter = create_adapter(config_for("yi")).expect("工厂应能创建 yi 适配器"); + assert_eq!(adapter.provider_type(), "yi"); + } + + #[test] + fn create_sensenova_returns_adapter() { + let adapter = + create_adapter(config_for("sensenova")).expect("工厂应能创建 sensenova 适配器"); + assert_eq!(adapter.provider_type(), "sensenova"); + } + + #[test] + fn create_hunyuan_returns_adapter() { + let adapter = create_adapter(config_for("hunyuan")).expect("工厂应能创建 hunyuan 适配器"); + assert_eq!(adapter.provider_type(), "hunyuan"); + } + + #[test] + fn create_groq_returns_adapter() { + let adapter = create_adapter(config_for("groq")).expect("工厂应能创建 groq 适配器"); + assert_eq!(adapter.provider_type(), "groq"); + } + + #[test] + fn create_deepseek_returns_adapter() { + let adapter = create_adapter(config_for("deepseek")).expect("工厂应能创建 deepseek 适配器"); + assert_eq!(adapter.provider_type(), "deepseek"); + } + + #[test] + fn create_stepfun_returns_adapter() { + let adapter = create_adapter(config_for("stepfun")).expect("工厂应能创建 stepfun 适配器"); + assert_eq!(adapter.provider_type(), "stepfun"); + } + + #[test] + fn create_mistral_returns_adapter() { + let adapter = create_adapter(config_for("mistral")).expect("工厂应能创建 mistral 适配器"); + assert_eq!(adapter.provider_type(), "mistral"); + } + + #[test] + fn create_cohere_returns_adapter() { + let adapter = create_adapter(config_for("cohere")).expect("工厂应能创建 cohere 适配器"); + assert_eq!(adapter.provider_type(), "cohere"); + } + + #[test] + fn create_perplexity_returns_adapter() { + let adapter = + create_adapter(config_for("perplexity")).expect("工厂应能创建 perplexity 适配器"); + assert_eq!(adapter.provider_type(), "perplexity"); + } + + #[test] + fn create_additional_models_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: + // xaigrok -> grok / lingyiwanwu -> yi / shangtang -> sensenova / tencent_hunyuan -> hunyuan + let xaigrok = create_adapter(config_for("xaigrok")).expect("别名 xaigrok 应映射到 grok"); + assert_eq!(xaigrok.provider_type(), "grok"); + let lingyiwanwu = + create_adapter(config_for("lingyiwanwu")).expect("别名 lingyiwanwu 应映射到 yi"); + assert_eq!(lingyiwanwu.provider_type(), "yi"); + let shangtang = + create_adapter(config_for("shangtang")).expect("别名 shangtang 应映射到 sensenova"); + assert_eq!(shangtang.provider_type(), "sensenova"); + let tencent_hunyuan = create_adapter(config_for("tencent_hunyuan")) + .expect("别名 tencent_hunyuan 应映射到 hunyuan"); + assert_eq!(tencent_hunyuan.provider_type(), "hunyuan"); + } + + #[test] + fn create_more_models_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/more_models.py 末尾 register 调用: + // step -> stepfun(deepseek/mistral/cohere/perplexity 无别名) + let step = create_adapter(config_for("step")).expect("别名 step 应映射到 stepfun"); + assert_eq!(step.provider_type(), "stepfun"); + } + #[test] fn create_phase2_adapter_returns_phase2_message() { let result = create_adapter(config_for("anthropic")); @@ -230,6 +346,17 @@ mod tests { assert!(is_known_provider("azure")); assert!(is_known_provider("siliconflow")); assert!(is_known_provider("cloudflareai")); + // 阶段 2a 第二批 additional_models + more_models + assert!(is_known_provider("grok")); + assert!(is_known_provider("yi")); + assert!(is_known_provider("sensenova")); + assert!(is_known_provider("hunyuan")); + assert!(is_known_provider("groq")); + assert!(is_known_provider("deepseek")); + assert!(is_known_provider("stepfun")); + assert!(is_known_provider("mistral")); + assert!(is_known_provider("cohere")); + assert!(is_known_provider("perplexity")); } #[test] diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 00f9914..0a7a1e5 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -33,3 +33,9 @@ pub mod azure; /// 聚合平台适配器:阶段 2a,含 SiliconFlow/TogetherAI/FireworksAI/CloudflareAI 四个 OpenAI 兼容子适配器 pub mod aggregation_platforms; + +/// 扩展模型适配器:阶段 2a,含 Grok/Yi/SenseNova/Hunyuan/Groq 五个 OpenAI 兼容子适配器 +pub mod additional_models; + +/// 更多模型适配器:阶段 2a,含 DeepSeek/StepFun/Mistral/Cohere/Perplexity 五个 OpenAI 兼容子适配器 +pub mod more_models; From ba4a78bfed9b753ce7dd78218dfd6385f2054030 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:37:00 +0800 Subject: [PATCH 29/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20emerging=5Fmodels=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../src/adapters/emerging_models.rs | 2723 +++++++++++++++++ 1 file changed, 2723 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/emerging_models.rs diff --git a/crates/aibridge-core/src/adapters/emerging_models.rs b/crates/aibridge-core/src/adapters/emerging_models.rs new file mode 100644 index 0000000..fac9969 --- /dev/null +++ b/crates/aibridge-core/src/adapters/emerging_models.rs @@ -0,0 +1,2723 @@ +//! 新兴模型适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/emerging_models.py`。 +//! +//! 支持三个 Provider(按协议分两类): +//! - **Ideogram**:独立协议(`Api-Key` header + `/generate` 端点 + `image_request` 包裹), +//! 文字渲染最强的图像生成平台。仅 image 能力。 +//! - **Luma Dream Machine**:独立协议(Bearer + `/generations` 端点),高质量视频生成。 +//! 仅 video 能力(video_create + video_poll)。 +//! - **Meta Llama**:OpenAI 兼容协议(`POST /chat/completions` + `GET /models`), +//! Meta 官方 Llama API。chat + chat_stream + vision 能力,全部委托 [`OpenAiCompatAdapter`] 地基。 +//! +//! ## 结构 +//! +//! Python v1 是 3 个独立 adapter 类,统一注册到工厂。Rust 同样实现 3 个独立 struct: +//! - [`IdeogramAdapter`]:独立协议,自带 HttpClient,实现 image_generate + list_models(硬编码) +//! - [`LumaAdapter`]:独立协议,自带 HttpClient,实现 video_create + video_poll + list_models(硬编码) +//! - [`LlamaAdapter`]:组合 [`OpenAiCompatAdapter`] 地基,chat/chat_stream/list_models 全部委托 +//! +//! ## provider_type 标识(对齐 Python v1) +//! +//! | Provider | provider_type | 别名 | +//! |---|---|---| +//! | Ideogram | `ideogram` | `ideo` | +//! | Luma Dream Machine | `luma` | `dream-machine` / `lumalabs` | +//! | Meta Llama | `llama` | `meta-llama` / `meta` | +//! +//! ## 阶段范围 +//! +//! 阶段 2a 实现各 Provider 的核心能力(Ideogram 图像、Luma 视频、Llama 对话)+ list_models。 +//! 不支持的能力(如 Ideogram 的 chat/video、Luma 的 chat/image、Llama 的 image/video) +//! 走 `Adapter` trait 默认实现返 `UnsupportedCapability`,与 Python v1 抛 +//! `UnsupportedCapabilityError` 行为一致。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::chat::{ChatCompletion, ChatRequest}; +use crate::model::common::{ModelInfo, ModelType, TaskStatus}; +use crate::model::image::{FileInput, ImageData, ImageRequest, ImageResult}; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +// ==================== 默认 Base URL ==================== + +/// Ideogram 默认 Base URL +/// +/// 对应 Python v1 `IdeogramAdapter.DEFAULT_BASE_URL`。 +pub const DEFAULT_IDEOGRAM_BASE_URL: &str = "https://api.ideogram.ai"; + +/// Luma Dream Machine 默认 Base URL +/// +/// 对应 Python v1 `LumaAdapter.DEFAULT_BASE_URL`(已含 `/dream-machine/v1` 前缀)。 +pub const DEFAULT_LUMA_BASE_URL: &str = "https://api.lumalabs.ai/dream-machine/v1"; + +/// Meta Llama 默认 Base URL +/// +/// 对应 Python v1 `LlamaAdapter.DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +pub const DEFAULT_LLAMA_BASE_URL: &str = "https://api.llama.com/v1"; + +// ==================== 能力集合构造 ==================== + +/// Ideogram 支持的能力集合 +/// +/// 对齐 Python v1 `IdeogramAdapter.supported_capabilities = ["image"]`。 +/// Rust 用 `ImageGenerate` 表达图像生成能力。 +fn ideogram_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ImageGenerate); + caps +} + +/// Luma 支持的能力集合 +/// +/// 对齐 Python v1 `LumaAdapter.supported_capabilities = ["video"]`。 +/// Rust 用 `VideoGenerate` 表达视频生成能力(含 video_create + video_poll)。 +fn luma_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps +} + +/// Meta Llama 支持的能力集合 +/// +/// 对齐 Python v1 `LlamaAdapter.supported_capabilities = ["chat", "vision"]`。 +/// chat_stream 虽未在 Python 显式声明,但 Python 实现了该方法且 OpenAI 兼容协议天然支持流式, +/// 故 Rust 一并声明 ChatStream(与 openai/azure 等兼容适配器保持一致)。 +fn llama_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +// ==================== Ideogram 图像生成适配器 ==================== + +/// Ideogram 适配器 +/// +/// 文字渲染最强的图像生成平台,支持 V2/V2A/V1 等模型。 +/// 官方 API 文档:https://developers.ideogram.com/ +/// +/// ## API 规范 +/// - Base URL: `https://api.ideogram.ai` +/// - 文生图: `POST /generate`(body 用 `image_request` 包裹) +/// - 图生图/Remix: `POST /remix`(body 用 `image_request` 包裹) +/// - 局部重绘: `POST /inpaint` +/// - 扩图: `POST /outpaint` +/// - 认证: `Api-Key` header(注意不是 Bearer Token) +/// +/// ## 阶段范围 +/// 阶段 2a 仅实现文生图(`/generate`)+ list_models(硬编码)。 +/// Remix/Inpaint/Outpaint 等图像编辑能力走 trait 默认实现返 UnsupportedCapability, +/// 待后续阶段补齐 `image_edit` trait 方法后再实现。 +pub struct IdeogramAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl IdeogramAdapter { + /// 创建 Ideogram 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_IDEOGRAM_BASE_URL`] 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_IDEOGRAM_BASE_URL.to_string()); + + // 构造 HttpClient:把 base_url 透传,便于 post_authed_json 等方法自动拼接相对路径 + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: ideogram_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .unwrap_or_else(|| DEFAULT_IDEOGRAM_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: ideogram_capabilities(), + } + } + + /// API key(可能为空,免费 provider 场景) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: ideogram)", cap.as_str()), + }) + } + } + + /// 发送带 `Api-Key` 认证的 POST JSON 请求,并用 Ideogram 错误映射处理响应 + /// + /// 注意:Ideogram 用 `Api-Key` header 而非 Bearer Token,与 OpenAI 兼容族不同。 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .header("Api-Key", self.api_key()) + .header("Content-Type", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 Ideogram `/generate` 请求体 + /// + /// 移植自 Python v1 `image_generate`: + /// - 内层 `image_request` 包含 prompt / model 及可选参数 + /// - aspect_ratio 归一化为大写(与 Python `aspect_ratio.upper()` 一致) + /// - num_images 上限 8(与 Python `min(int(num_images), 8)` 一致) + /// - magic_prompt_level 映射到 `magic_prompt_option` 字段(与 Python 一致) + /// - extra 字段合并到内层 image_request(透传厂商特有参数) + fn build_generate_body(&self, req: &ImageRequest) -> Value { + // 默认模型 V_2A_TURBO(与 Python 一致) + let model = if req.model.is_empty() { + "V_2A_TURBO".to_string() + } else { + req.model.clone() + }; + + let mut image_request = json!({ + "prompt": req.prompt, + "model": model, + }); + + // 负面提示词 + if let Some(np) = &req.negative_prompt { + image_request["negative_prompt"] = json!(np); + } + // 宽高比(归一化为大写) + if let Some(ar) = &req.aspect_ratio { + image_request["aspect_ratio"] = json!(ar.to_uppercase()); + } + // 分辨率(Ideogram 用 resolution 字段表示分辨率,如 "1536x1536") + if let Some(size) = &req.size { + image_request["resolution"] = json!(size); + } + // 风格类型(style 字段映射到 style_type) + if let Some(style) = &req.style { + image_request["style_type"] = json!(style); + } + // 魔法提示词增强:Python 用 kwargs.magic_prompt_level 映射到 magic_prompt_option + if let Some(mp) = req.extra.get("magic_prompt_level") { + image_request["magic_prompt_option"] = mp.clone(); + } + // 生成数量(上限 8) + let n = req.n.min(8); + image_request["num_images"] = json!(n); + // 随机种子 + if let Some(seed) = req.seed { + image_request["seed"] = json!(seed); + } + + // extra 透传(合并到内层 image_request,跳过已处理的 magic_prompt_level) + if let Some(obj) = image_request.as_object_mut() { + for (k, v) in &req.extra { + if k != "magic_prompt_level" { + obj.insert(k.clone(), v.clone()); + } + } + } + + // Ideogram 用 image_request 包裹 + json!({ "image_request": image_request }) + } + + /// 解析 Ideogram `/generate` 响应 → ImageResult + /// + /// Ideogram 响应结构: + /// ```json + /// { + /// "request_id": "...", + /// "data": [{"url": "...", "base64": "...", "prompt": "..."}] + /// } + /// ``` + /// 注意:Ideogram 用 `base64` 而非 `b64_json`,`prompt` 而非 `revised_prompt`。 + fn parse_image_result(value: &Value, fallback_model: &str) -> Result { + let id = value + .get("request_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("img")); + let created = value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + let model = fallback_model.to_string(); + let data = value + .get("data") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().map(parse_ideogram_image_data).collect()) + .unwrap_or_default(); + Ok(ImageResult { + id, + object: "image.generation".to_string(), + created, + model, + data, + }) + } + + /// 将 Ideogram API 错误响应映射为 AibridgeError + /// + /// 移植自 Python v1 `_handle_ideogram_error`: + /// - 401 → Authentication("Invalid Ideogram API key") + /// - 402 → Api("Ideogram payment required or credits exhausted",额度耗尽) + /// - 429 → RateLimit("Ideogram rate limit exceeded or credits exhausted") + /// - 其余 ≥400 → Api(提取 message / error / detail,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Ideogram API key".to_string(), + }, + 402 => AibridgeError::Api { + status, + message: "Ideogram payment required or credits exhausted".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Ideogram rate limit exceeded or credits exhausted".to_string(), + retry_after: None, + }, + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +/// 解析单个 Ideogram 图像数据项 +/// +/// Ideogram 字段映射(与 OpenAI 不同): +/// - `url` → `url` +/// - `base64` → `b64_json` +/// - `prompt` → `revised_prompt`(模型优化后的提示词) +fn parse_ideogram_image_data(v: &Value) -> ImageData { + ImageData { + url: v.get("url").and_then(|x| x.as_str()).map(str::to_owned), + b64_json: v.get("base64").and_then(|x| x.as_str()).map(str::to_owned), + revised_prompt: v.get("prompt").and_then(|x| x.as_str()).map(str::to_owned), + } +} + +#[async_trait] +impl Adapter for IdeogramAdapter { + fn provider_type(&self) -> &str { + "ideogram" + } + + fn provider_name(&self) -> &str { + "Ideogram" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 图像生成:`POST /generate`(body 用 `image_request` 包裹) + async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let body = self.build_generate_body(&req); + let model = if req.model.is_empty() { + "V_2A_TURBO".to_string() + } else { + req.model.clone() + }; + let value = self.post_authed_json("generate", &body).await?; + Self::parse_image_result(&value, &model) + } + + /// 模型列表(硬编码) + /// + /// Ideogram 无标准 `/models` 端点,暂保留硬编码列表(与 Python v1 一致)。 + /// 含 V_2A / V_2A_TURBO / V_2 / V_1 / V_1_TURBO 五个模型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = ideogram_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / video / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== Luma Dream Machine 视频生成适配器 ==================== + +/// Luma Dream Machine 适配器 +/// +/// 高质量视频生成平台,支持文生视频和图生视频。 +/// 官方 API 文档:https://docs.lumalabs.ai/ +/// +/// ## API 规范 +/// - Base URL: `https://api.lumalabs.ai/dream-machine/v1` +/// - 创建生成: `POST /generations` +/// - 查询状态: `GET /generations/{id}` +/// - 认证: `Authorization: Bearer {api_key}` +/// +/// ## 阶段范围 +/// 阶段 2a 实现 video_create + video_poll + list_models(硬编码)。 +pub struct LumaAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl LumaAdapter { + /// 创建 Luma 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_LUMA_BASE_URL`] 兜底。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_LUMA_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: luma_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .unwrap_or_else(|| DEFAULT_LUMA_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: luma_capabilities(), + } + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: luma)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用 Luma 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 Bearer 认证的 GET 请求,并用 Luma 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 Luma `/generations` 请求体 + /// + /// 移植自 Python v1 `video_create`: + /// - prompt / model 必选 + /// - aspect_ratio / duration / resolution / loop / negative_prompt / camera_motion 透传 + /// - reference_images[0] → keyframes.frame0(图生视频起始帧) + /// - first_frame → keyframes.frame0(优先于 reference_images) + /// - last_frame → keyframes.frame1(结束帧) + /// - extra 字段合并到顶层(透传厂商特有参数) + fn build_generations_body(&self, req: &VideoRequest) -> Value { + // 默认模型 ray-2(与 Python 一致) + let model = if req.model.is_empty() { + "ray-2".to_string() + } else { + req.model.clone() + }; + + let mut body = json!({ + "prompt": req.prompt, + "model": model, + }); + + // 宽高比 + if let Some(ar) = &req.aspect_ratio { + body["aspect_ratio"] = json!(ar); + } + // 时长(VideoRequest.duration 是 u32 秒数,Luma 接受 "5s"/"9s" 字符串) + if let Some(d) = req.duration { + body["duration"] = json!(format!("{d}s")); + } + // 分辨率 + if let Some(r) = &req.resolution { + body["resolution"] = json!(r); + } + // 循环(with_audio 字段不语义对应 loop,用 extra 里的 loop 优先) + if let Some(lp) = req.extra.get("loop") { + body["loop"] = lp.clone(); + } + // 负面提示词 + if let Some(np) = &req.negative_prompt { + body["negative_prompt"] = json!(np); + } + // 相机运动 + if let Some(cm) = &req.camera_motion { + body["camera_motion"] = json!(cm); + } + + // 关键帧:first_frame / last_frame / reference_images[0] → keyframes + let mut keyframes = json!({}); + // first_frame 优先作为 frame0 + if let Some(ff) = &req.first_frame { + keyframes["frame0"] = json!({ "type": "image", "url": file_input_to_url(ff) }); + } else if !req.reference_images.is_empty() { + // reference_images[0] 作为 frame0(图生视频起始帧) + keyframes["frame0"] = + json!({ "type": "image", "url": file_input_to_url(&req.reference_images[0]) }); + } + // last_frame 作为 frame1(结束帧) + if let Some(lf) = &req.last_frame { + keyframes["frame1"] = json!({ "type": "image", "url": file_input_to_url(lf) }); + } + // extra 里的 keyframes 直接覆盖(更细粒度控制) + if let Some(custom_kf) = req.extra.get("keyframes") { + if let Some(custom_obj) = custom_kf.as_object() { + if let Some(obj) = keyframes.as_object_mut() { + for (k, v) in custom_obj { + obj.insert(k.clone(), v.clone()); + } + } + } + } + if !keyframes.as_object().map(|o| o.is_empty()).unwrap_or(true) { + body["keyframes"] = keyframes; + } + + // extra 透传(合并到顶层,跳过已处理的 loop / keyframes) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + if k != "loop" && k != "keyframes" { + obj.insert(k.clone(), v.clone()); + } + } + } + + body + } + + /// 解析 Luma `/generations` 创建响应 → VideoTask + /// + /// Luma 创建响应:`{"id": "...", "state": "queued", ...}` + fn parse_video_task(value: &Value, model: &str) -> Result { + let task_id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("vid")); + let state = value + .get("state") + .and_then(|v| v.as_str()) + .unwrap_or("queued"); + Ok(VideoTask { + task_id, + model: model.to_string(), + status: map_luma_status(state), + created_at: util::current_timestamp(), + }) + } + + /// 解析 Luma `/generations/{id}` 查询响应 → VideoStatus + /// + /// 移植自 Python v1 `video_poll`: + /// - 视频 URL 在 `assets.video` / `assets.mp4` / `video` + /// - 错误信息在 `failure_reason` / `error` + /// - 进度估算:success=100, dreaming=30, processing=70, queued/pending=5 + /// - created_at / updated_at 为 ISO 时间字符串,转时间戳 + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let state = value.get("state").and_then(|v| v.as_str()).unwrap_or(""); + let status = map_luma_status(state); + + // 视频 URL:assets.video 优先,兼容 assets.mp4 / video + let video_url = if status == TaskStatus::Success { + value + .get("assets") + .and_then(|a| a.get("video")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("assets") + .and_then(|a| a.get("mp4")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("video") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + } else { + None + }; + + // 错误信息 + let error = if status == TaskStatus::Failed { + value + .get("failure_reason") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| Some("Generation failed".to_string())) + } else { + None + }; + + // 进度估算(先判断 state,因为 dreaming/processing 都映射到 processing 但进度不同) + let state_lower = state.to_lowercase(); + let progress = if status == TaskStatus::Success { + 100 + } else if state_lower == "dreaming" { + 30 + } else if status == TaskStatus::Processing { + 70 + } else if state_lower == "queued" || state_lower == "pending" { + 5 + } else { + 0 + }; + + // Luma 返回 ISO 时间字符串,转为时间戳 + let created_at = value + .get("created_at") + .and_then(|v| v.as_str()) + .and_then(parse_iso_to_timestamp); + let updated_at = value + .get("updated_at") + .and_then(|v| v.as_str()) + .and_then(parse_iso_to_timestamp) + .or_else(|| Some(util::current_timestamp())); + + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress: Some(progress), + error, + created_at, + updated_at, + } + } + + /// 将 Luma API 错误响应映射为 AibridgeError + /// + /// 移植自 Python v1 `_handle_luma_error`: + /// - 401 → Authentication("Invalid Luma API key") + /// - 402 → Api("Luma credits exhausted or payment required") + /// - 429 → RateLimit("Luma rate limit exceeded or credits exhausted") + /// - 其余 ≥400 → Api(提取 detail / error.message / message,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Luma API key".to_string(), + }, + 402 => AibridgeError::Api { + status, + message: "Luma credits exhausted or payment required".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Luma rate limit exceeded or credits exhausted".to_string(), + retry_after: None, + }, + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for LumaAdapter { + fn provider_type(&self) -> &str { + "luma" + } + + fn provider_name(&self) -> &str { + "Luma Dream Machine" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 创建视频生成任务:`POST /generations` + async fn video_create(&self, req: VideoRequest) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let body = self.build_generations_body(&req); + let model = if req.model.is_empty() { + "ray-2".to_string() + } else { + req.model.clone() + }; + let value = self.post_authed_json("generations", &body).await?; + Self::parse_video_task(&value, &model) + } + + /// 查询视频任务状态:`GET /generations/{id}` + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let path = format!("generations/{task_id}"); + let value = self.get_authed_json(&path).await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 模型列表(硬编码) + /// + /// Luma 无标准 `/models` 端点,暂保留硬编码列表(与 Python v1 一致)。 + /// 含 ray-2 / ray-2-flash / dream-machine 三个模型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = luma_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / image / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== Meta Llama 适配器(OpenAI 兼容)==================== + +/// Meta Llama 适配器 +/// +/// Meta 官方 Llama API,完全兼容 OpenAI 接口规范。 +/// 官方 API 文档:https://docs.llama.com/ +/// +/// ## API 规范 +/// - Base URL: `https://api.llama.com/v1` +/// - Chat: `POST /chat/completions`(OpenAI 兼容) +/// - Models: `GET /models`(OpenAI 兼容) +/// - 认证: Bearer Token +/// - 支持流式输出 +/// +/// 全部能力(chat / chat_stream / list_models)委托给 [`OpenAiCompatAdapter`] 地基, +/// 仅 base_url / provider_type / capabilities 差异。 +pub struct LlamaAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl LlamaAdapter { + /// 创建 Meta Llama 适配器 + /// + /// `config.base_url` 为空时回退到 [`DEFAULT_LLAMA_BASE_URL`]。 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + "llama", + "Meta Llama", + DEFAULT_LLAMA_BASE_URL, + llama_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 地基构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } +} + +#[async_trait] +impl Adapter for LlamaAdapter { + fn provider_type(&self) -> &str { + "llama" + } + + fn provider_name(&self) -> &str { + "Meta Llama" + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取) + /// + /// 调用 `GET /models`(OpenAI 兼容端点),实时拉取模型列表。 + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 内部:硬编码模型列表 ==================== + +/// Ideogram 硬编码模型列表 +/// +/// 对应 Python v1 `IdeogramAdapter.list_models`。 +/// 注意:该 Provider 无标准 `/models` 端点,暂保留硬编码列表。 +fn ideogram_hardcoded_models() -> Vec { + vec![ + ModelInfo { + id: "V_2A".into(), + name: "Ideogram V2A".into(), + model_type: ModelType::Image, + provider: "ideogram".into(), + capabilities: vec!["text2image".into(), "image2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Ideogram V2A 文生图模型,文字渲染强".into()), + created: None, + }, + ModelInfo { + id: "V_2A_TURBO".into(), + name: "Ideogram V2A Turbo".into(), + model_type: ModelType::Image, + provider: "ideogram".into(), + capabilities: vec!["text2image".into(), "image2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Ideogram V2A Turbo 快速版本,文字渲染强".into()), + created: None, + }, + ModelInfo { + id: "V_2".into(), + name: "Ideogram V2".into(), + model_type: ModelType::Image, + provider: "ideogram".into(), + capabilities: vec!["text2image".into(), "image2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Ideogram V2 高质量模型".into()), + created: None, + }, + ModelInfo { + id: "V_1".into(), + name: "Ideogram V1".into(), + model_type: ModelType::Image, + provider: "ideogram".into(), + capabilities: vec!["text2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Ideogram V1 标准模型".into()), + created: None, + }, + ModelInfo { + id: "V_1_TURBO".into(), + name: "Ideogram V1 Turbo".into(), + model_type: ModelType::Image, + provider: "ideogram".into(), + capabilities: vec!["text2image".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Ideogram V1 Turbo 快速模型".into()), + created: None, + }, + ] +} + +/// Luma 硬编码模型列表 +/// +/// 对应 Python v1 `LumaAdapter.list_models`。 +/// 注意:该 Provider 无标准 `/models` 端点,暂保留硬编码列表。 +fn luma_hardcoded_models() -> Vec { + vec![ + ModelInfo { + id: "ray-2".into(), + name: "Luma Ray 2".into(), + model_type: ModelType::Video, + provider: "luma".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Luma Ray 2 高质量视频生成模型".into()), + created: None, + }, + ModelInfo { + id: "ray-2-flash".into(), + name: "Luma Ray 2 Flash".into(), + model_type: ModelType::Video, + provider: "luma".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Luma Ray 2 Flash 快速视频生成模型".into()), + created: None, + }, + ModelInfo { + id: "dream-machine".into(), + name: "Dream Machine".into(), + model_type: ModelType::Video, + provider: "luma".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Luma Dream Machine 初代视频模型".into()), + created: None, + }, + ] +} + +// ==================== 内部:辅助函数 ==================== + +/// 映射 Luma 状态到标准 TaskStatus +/// +/// 移植自 Python v1 `_map_luma_status`: +/// - queued / pending → Pending +/// - dreaming / processing → Processing +/// - completed / succeeded / success → Success +/// - failed / error → Failed +/// - 未知 → Pending(与 Python 默认值一致) +fn map_luma_status(raw_state: &str) -> TaskStatus { + match raw_state.to_lowercase().as_str() { + "queued" | "pending" => TaskStatus::Pending, + "dreaming" | "processing" => TaskStatus::Processing, + "completed" | "succeeded" | "success" => TaskStatus::Success, + "failed" | "error" => TaskStatus::Failed, + _ => TaskStatus::Pending, + } +} + +/// 将 FileInput 转为 URL 字符串 +/// +/// Luma keyframes 需要 URL 形式的图像引用: +/// - `Url(s)` → 直接返回 s +/// - `Base64(s)` → 返回 s(Luma 接受 base64 字符串作为 url) +/// - `Path(_)` / `Bytes(_)` → 空字符串(这些形式需上层预转换为 URL/base64,此处不处理) +/// +/// 与 Python v1 行为一致:Python 仅处理 data: 前缀和 http URL,其余原样透传。 +fn file_input_to_url(input: &FileInput) -> String { + match input { + FileInput::Url(s) | FileInput::Base64(s) => s.clone(), + FileInput::Path(_) | FileInput::Bytes(_) => String::new(), + } +} + +/// 解析 ISO 8601 时间字符串为 Unix 时间戳 +/// +/// Luma 返回 ISO 时间字符串(如 `2024-01-01T00:00:00Z`),转为时间戳。 +/// 解析失败返回 None(与 Python try/except 行为一致)。 +fn parse_iso_to_timestamp(s: &str) -> Option { + // 处理 Z 后缀(替换为 +00:00 便于统一解析) + let normalized = if let Some(stripped) = s.strip_suffix('Z') { + format!("{stripped}+00:00") + } else { + s.to_string() + }; + parse_rfc3339(&normalized) +} + +/// 手动解析 RFC3339 时间字符串为 Unix 时间戳 +/// +/// 支持格式:`YYYY-MM-DDTHH:MM:SS[.fff][+HH:MM | -HH:MM | Z]` +/// 解析失败返回 None。不引入 chrono 依赖,用 Howard Hinnant 的 civil_from_days 算法。 +fn parse_rfc3339(s: &str) -> Option { + // 格式:YYYY-MM-DDTHH:MM:SS + if s.len() < 19 { + return None; + } + let bytes = s.as_bytes(); + // 解析日期时间部分 YYYY-MM-DDTHH:MM:SS + let year: u32 = std::str::from_utf8(&bytes[0..4]).ok()?.parse().ok()?; + if bytes[4] != b'-' + || bytes[7] != b'-' + || bytes[10] != b'T' + || bytes[13] != b':' + || bytes[16] != b':' + { + return None; + } + let month: u32 = std::str::from_utf8(&bytes[5..7]).ok()?.parse().ok()?; + let day: u32 = std::str::from_utf8(&bytes[8..10]).ok()?.parse().ok()?; + let hour: u32 = std::str::from_utf8(&bytes[11..13]).ok()?.parse().ok()?; + let minute: u32 = std::str::from_utf8(&bytes[14..16]).ok()?.parse().ok()?; + let second: u32 = std::str::from_utf8(&bytes[17..19]).ok()?.parse().ok()?; + + // 解析时区偏移:[+|-]HH:MM 或 Z(已规范化为 +00:00) + let tz_offset_seconds: i64 = if s.len() >= 25 { + let sign = match bytes[19] { + b'+' => 1i64, + b'-' => -1i64, + _ => return None, + }; + let tz_hour: i64 = std::str::from_utf8(&bytes[20..22]).ok()?.parse().ok()?; + let tz_minute: i64 = std::str::from_utf8(&bytes[23..25]).ok()?.parse().ok()?; + sign * (tz_hour * 3600 + tz_minute * 60) + } else { + // 无时区信息,按 UTC 处理 + 0 + }; + + // 转为 Unix 时间戳(UTC) + let utc_seconds = civil_to_unix(year, month, day, hour, minute, second)?; + // 减去时区偏移得到 UTC 时间戳 + Some((utc_seconds as i64 - tz_offset_seconds) as u64) +} + +/// 公历日期时间转 Unix 时间戳(UTC) +/// +/// 算法:Howard Hinnant 的 civil_from_days,从 1970-01-01 起累加天数。 +/// 返回 None 表示日期非法或超出支持范围。 +fn civil_to_unix( + year: u32, + month: u32, + day: u32, + hour: u32, + minute: u32, + second: u32, +) -> Option { + if !(1..=12).contains(&month) + || !(1..=31).contains(&day) + || hour > 23 + || minute > 59 + || second > 59 + { + return None; + } + let y = year as i64; + let m = month as i64; + let d = day as i64; + // 调整:3月为年初(避免闰年判断的边界问题) + let y_adj = if m <= 2 { y - 1 } else { y }; + let era = if y_adj >= 0 { y_adj } else { y_adj - 399 } / 400; + let yoe = (y_adj - era * 400) as u64; // [0, 399] + let doy = (153 * (if m > 2 { m - 3 } else { m + 9 }) + 2) / 5 + d - 1; // [0, 365] + let doe = yoe * 365 + yoe / 4 - yoe / 100 + doy as u64; // [0, 146096] + let days = era * 146097 + doe as i64 - 719468; + let seconds = days * 86400 + (hour as i64) * 3600 + (minute as i64) * 60 + second as i64; + Some(seconds) +} + +/// 解析错误体中的 message 字段 +/// +/// 通用错误消息提取,兼容多种错误体结构: +/// - `{"error": {"message": "..."}}`(OpenAI 风格) +/// - `{"message": "..."}`(部分平台顶层 message) +/// - `{"detail": "..."}`(Luma 用 detail 字段) +/// - `{"detail": ["err1", "err2"]}`(Luma validation 错误数组) +/// - `{"error": "..."}`(Ideogram 用顶层 error 字符串) +/// +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + // Luma 优先用 detail 字段 + if let Some(msg) = v.get("detail").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // detail 可能是数组(Luma 的 validation 错误) + if let Some(arr) = v.get("detail").and_then(|m| m.as_array()) { + let parts: Vec = arr + .iter() + .filter_map(|x| x.as_str().map(str::to_owned)) + .collect(); + if !parts.is_empty() { + return parts.join("; "); + } + } + // OpenAI 风格 error.message + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 顶层 error 字符串(Ideogram 风格) + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::error::AibridgeError; + use crate::http::HttpClient; + use crate::model::chat::ChatMessage; + use crate::model::image::FileInput; + use futures::stream::StreamExt; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 IdeogramAdapter(指向 mockito server) + fn make_ideogram(server: &Server) -> IdeogramAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("ideogram", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + IdeogramAdapter::with_http(http, config) + } + + /// 构造测试用 LumaAdapter(指向 mockito server) + fn make_luma(server: &Server) -> LumaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("luma", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + LumaAdapter::with_http(http, config) + } + + /// 构造测试用 LlamaAdapter(指向 mockito server) + fn make_llama(server: &Server) -> LlamaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("llama", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let compat = OpenAiCompatAdapter::with_http( + http, + config, + "llama", + "Meta Llama", + llama_capabilities(), + ); + LlamaAdapter::with_compat(compat) + } + + /// 构造不指向任何 server 的 IdeogramAdapter(用于不发请求的元信息/能力测试) + fn make_ideogram_no_server() -> IdeogramAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_IDEOGRAM_BASE_URL) + .build(); + let config = ProviderConfig::from_options("ideogram", opts); + IdeogramAdapter::new(config).expect("IdeogramAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 LumaAdapter + fn make_luma_no_server() -> LumaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_LUMA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("luma", opts); + LumaAdapter::new(config).expect("LumaAdapter 构造应成功") + } + + /// 构造不指向任何 server 的 LlamaAdapter + fn make_llama_no_server() -> LlamaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_LLAMA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("llama", opts); + LlamaAdapter::new(config).expect("LlamaAdapter 构造应成功") + } + + // ============ Ideogram 元信息 ============ + + #[test] + fn ideogram_provider_type_and_name_match_python() { + let adapter = make_ideogram_no_server(); + assert_eq!(adapter.provider_type(), "ideogram"); + assert_eq!(adapter.provider_name(), "Ideogram"); + } + + #[test] + fn ideogram_requires_api_key_is_true() { + let adapter = make_ideogram_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn ideogram_capabilities_contains_only_image() { + let adapter = make_ideogram_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::ImageGenerate)); + // chat / video 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::VideoGenerate)); + } + + #[test] + fn ideogram_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("ideogram", opts); + let adapter = IdeogramAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_IDEOGRAM_BASE_URL); + } + + #[test] + fn ideogram_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.ideogram-proxy.com") + .build(); + let config = ProviderConfig::from_options("ideogram", opts); + let adapter = IdeogramAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.ideogram-proxy.com"); + } + + // ============ Ideogram image_generate ============ + + #[tokio::test] + async fn ideogram_image_generate_success_parses_url() { + let mut server = Server::new_async().await; + let body = json!({ + "request_id": "req-123", + "data": [{ + "url": "https://example.com/img.png", + "prompt": "a cute cat" + }] + }); + let mock = server + .mock("POST", "/generate") + .match_header("Api-Key", "test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat") + .aspect_ratio("16:9") + .n(2) + .seed(42) + .build(); + let resp = adapter + .image_generate(req) + .await + .expect("image_generate 应成功"); + + assert_eq!(resp.id, "req-123"); + assert_eq!(resp.model, "V_2A"); + assert_eq!(resp.data.len(), 1); + assert_eq!( + resp.data[0].url.as_deref(), + Some("https://example.com/img.png") + ); + assert_eq!(resp.data[0].revised_prompt.as_deref(), Some("a cute cat")); + mock.assert_async().await; + } + + #[tokio::test] + async fn ideogram_image_generate_parses_base64() { + let mut server = Server::new_async().await; + let body = json!({ + "request_id": "req-456", + "data": [{"base64": "aGVsbG8="}] + }); + server + .mock("POST", "/generate") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A_TURBO", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + } + + #[tokio::test] + async fn ideogram_image_generate_wraps_in_image_request() { + // 验证请求体用 image_request 包裹,且 aspect_ratio 归一化为大写 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generate") + .match_body(mockito::Matcher::PartialJson(json!({ + "image_request": { + "model": "V_2A", + "prompt": "a cat", + "aspect_ratio": "16:9", + "num_images": 3 + } + }))) + .with_status(200) + .with_body(json!({"request_id": "x", "data": []}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat") + .aspect_ratio("16:9") + .n(3) + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn ideogram_image_generate_uses_default_model_when_empty() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generate") + .match_body(mockito::Matcher::PartialJson(json!({ + "image_request": {"model": "V_2A_TURBO"} + }))) + .with_status(200) + .with_body(json!({"request_id": "x", "data": []}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.model, "V_2A_TURBO"); + mock.assert_async().await; + } + + #[tokio::test] + async fn ideogram_image_generate_caps_num_images_at_8() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generate") + .match_body(mockito::Matcher::PartialJson(json!({ + "image_request": {"num_images": 8} + }))) + .with_status(200) + .with_body(json!({"data": []}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + // 请求 20 张,应被截断为 8 + let req = ImageRequest::builder("V_2A", "a cat").n(20).build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn ideogram_image_generate_passes_extra_params() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generate") + .match_body(mockito::Matcher::PartialJson(json!({ + "image_request": { + "magic_prompt_option": "HIGH", + "style_type": "REALISTIC", + "custom_param": "custom_value" + } + }))) + .with_status(200) + .with_body(json!({"data": []}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat") + .style("REALISTIC") + .extra("magic_prompt_level", "HIGH") + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + // ============ Ideogram image_generate 错误路径 ============ + + #[tokio::test] + async fn ideogram_image_generate_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generate") + .with_status(401) + .with_body(json!({"error": "Invalid API key"}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn ideogram_image_generate_error_402_returns_api_payment() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generate") + .with_status(402) + .with_body(json!({"message": "credits exhausted"}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 402); + assert!(message.contains("credits exhausted")); + } + _ => panic!("应为 Api (402)"), + } + } + + #[tokio::test] + async fn ideogram_image_generate_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generate") + .with_status(429) + .with_body(json!({"error": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn ideogram_image_generate_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generate") + .with_status(500) + .with_body(json!({"error": "internal"}).to_string()) + .create_async() + .await; + + let adapter = make_ideogram(&server); + let req = ImageRequest::builder("V_2A", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn ideogram_chat_returns_unsupported() { + let adapter = make_ideogram_no_server(); + let req = ChatRequest::builder("V_2A", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn ideogram_video_create_returns_unsupported() { + let adapter = make_ideogram_no_server(); + let req = VideoRequest::builder("V_2A", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ Ideogram list_models ============ + + #[tokio::test] + async fn ideogram_list_models_returns_hardcoded() { + let adapter = make_ideogram_no_server(); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 5); + assert_eq!(models[0].id, "V_2A"); + assert_eq!(models[0].provider, "ideogram"); + assert_eq!(models[0].model_type, ModelType::Image); + } + + #[tokio::test] + async fn ideogram_list_models_filter_by_image_type() { + let adapter = make_ideogram_no_server(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert_eq!(images.len(), 5); + assert!(images.iter().all(|m| m.model_type == ModelType::Image)); + } + + #[tokio::test] + async fn ideogram_list_models_filter_by_video_returns_empty() { + let adapter = make_ideogram_no_server(); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert!(videos.is_empty()); + } + + // ============ Luma 元信息 ============ + + #[test] + fn luma_provider_type_and_name_match_python() { + let adapter = make_luma_no_server(); + assert_eq!(adapter.provider_type(), "luma"); + assert_eq!(adapter.provider_name(), "Luma Dream Machine"); + } + + #[test] + fn luma_requires_api_key_is_true() { + let adapter = make_luma_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn luma_capabilities_contains_only_video() { + let adapter = make_luma_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + // chat / image 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn luma_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("luma", opts); + let adapter = LumaAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_LUMA_BASE_URL); + } + + #[test] + fn luma_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.luma-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("luma", opts); + let adapter = LumaAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.luma-proxy.com/v1"); + } + + // ============ Luma video_create ============ + + #[tokio::test] + async fn luma_video_create_success_returns_task() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "state": "queued" + }); + let mock = server + .mock("POST", "/generations") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat running") + .aspect_ratio("16:9") + .duration(5) + .resolution("720p") + .build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + + assert_eq!(task.task_id, "gen-abc"); + assert_eq!(task.model, "ray-2"); + assert_eq!(task.status, TaskStatus::Pending); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_create_sends_model_and_prompt() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "ray-2", + "prompt": "a cat", + "aspect_ratio": "16:9", + "duration": "5s", + "resolution": "720p" + }))) + .with_status(200) + .with_body(json!({"id": "x", "state": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat") + .aspect_ratio("16:9") + .duration(5) + .resolution("720p") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_create_uses_default_model_when_empty() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "ray-2" + }))) + .with_status(200) + .with_body(json!({"id": "x", "state": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.model, "ray-2"); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_create_with_first_frame_builds_keyframes() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "keyframes": { + "frame0": {"type": "image", "url": "https://example.com/start.png"}, + "frame1": {"type": "image", "url": "https://example.com/end.png"} + } + }))) + .with_status(200) + .with_body(json!({"id": "x", "state": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat") + .first_frame(FileInput::url("https://example.com/start.png")) + .last_frame(FileInput::url("https://example.com/end.png")) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_create_with_reference_images_builds_frame0() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "keyframes": { + "frame0": {"type": "image", "url": "https://example.com/ref.png"} + } + }))) + .with_status(200) + .with_body(json!({"id": "x", "state": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat") + .reference_images(vec![FileInput::url("https://example.com/ref.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_create_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(401) + .with_body(json!({"detail": "Invalid API key"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn luma_video_create_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(429) + .with_body(json!({"detail": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn luma_video_create_error_402_returns_api_payment() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(402) + .with_body(json!({"detail": "credits exhausted"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 402), + _ => panic!("应为 Api (402)"), + } + } + + #[tokio::test] + async fn luma_video_create_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(500) + .with_body(json!({"detail": "internal"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let req = VideoRequest::builder("ray-2", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ Luma video_poll ============ + + #[tokio::test] + async fn luma_video_poll_success_returns_video_url() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "state": "completed", + "assets": {"video": "https://example.com/video.mp4"}, + "created_at": "2024-01-01T00:00:00Z", + "updated_at": "2024-01-01T00:00:05Z" + }); + let mock = server + .mock("GET", "/generations/gen-abc") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter + .video_poll("gen-abc", "ray-2") + .await + .expect("video_poll 应成功"); + + assert_eq!(status.task_id, "gen-abc"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/video.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert!(status.error.is_none()); + assert!(status.created_at.is_some()); + assert!(status.updated_at.is_some()); + mock.assert_async().await; + } + + #[tokio::test] + async fn luma_video_poll_processing_returns_progress() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "state": "dreaming" + }); + server + .mock("GET", "/generations/gen-abc") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter.video_poll("gen-abc", "ray-2").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(30)); + assert!(status.video_url.is_none()); + } + + #[tokio::test] + async fn luma_video_poll_failed_returns_error() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "state": "failed", + "failure_reason": "content policy violation" + }); + server + .mock("GET", "/generations/gen-abc") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter.video_poll("gen-abc", "ray-2").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("content policy violation")); + } + + #[tokio::test] + async fn luma_video_poll_processing_state_returns_70_progress() { + let mut server = Server::new_async().await; + let body = json!({"id": "x", "state": "processing"}); + server + .mock("GET", "/generations/x") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter.video_poll("x", "ray-2").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(70)); + } + + #[tokio::test] + async fn luma_video_poll_queued_returns_5_progress() { + let mut server = Server::new_async().await; + let body = json!({"id": "x", "state": "queued"}); + server + .mock("GET", "/generations/x") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter.video_poll("x", "ray-2").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + assert_eq!(status.progress, Some(5)); + } + + #[tokio::test] + async fn luma_video_poll_assets_mp4_fallback() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "x", + "state": "completed", + "assets": {"mp4": "https://example.com/video.mp4"} + }); + server + .mock("GET", "/generations/x") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let status = adapter.video_poll("x", "ray-2").await.unwrap(); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/video.mp4") + ); + } + + #[tokio::test] + async fn luma_video_poll_error_404_returns_api() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/nonexistent") + .with_status(404) + .with_body(json!({"detail": "not found"}).to_string()) + .create_async() + .await; + + let adapter = make_luma(&server); + let err = adapter + .video_poll("nonexistent", "ray-2") + .await + .unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 404), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn luma_chat_returns_unsupported() { + let adapter = make_luma_no_server(); + let req = ChatRequest::builder("ray-2", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn luma_image_generate_returns_unsupported() { + let adapter = make_luma_no_server(); + let req = ImageRequest::builder("ray-2", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ Luma list_models ============ + + #[tokio::test] + async fn luma_list_models_returns_hardcoded() { + let adapter = make_luma_no_server(); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + assert_eq!(models[0].id, "ray-2"); + assert_eq!(models[0].provider, "luma"); + assert_eq!(models[0].model_type, ModelType::Video); + } + + #[tokio::test] + async fn luma_list_models_filter_by_video_type() { + let adapter = make_luma_no_server(); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert_eq!(videos.len(), 3); + assert!(videos.iter().all(|m| m.model_type == ModelType::Video)); + } + + #[tokio::test] + async fn luma_list_models_filter_by_image_returns_empty() { + let adapter = make_luma_no_server(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert!(images.is_empty()); + } + + // ============ Llama 元信息 ============ + + #[test] + fn llama_provider_type_and_name_match_python() { + let adapter = make_llama_no_server(); + assert_eq!(adapter.provider_type(), "llama"); + assert_eq!(adapter.provider_name(), "Meta Llama"); + } + + #[test] + fn llama_requires_api_key_is_true() { + let adapter = make_llama_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn llama_capabilities_contains_chat_and_vision() { + let adapter = make_llama_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // image / video 不声明 + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::VideoGenerate)); + } + + #[test] + fn llama_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("llama", opts); + let adapter = LlamaAdapter::new(config).unwrap(); + assert_eq!(adapter.compat.base_url(), DEFAULT_LLAMA_BASE_URL); + } + + #[test] + fn llama_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.llama-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("llama", opts); + let adapter = LlamaAdapter::new(config).unwrap(); + assert_eq!( + adapter.compat.base_url(), + "https://custom.llama-proxy.com/v1" + ); + } + + // ============ Llama chat ============ + + #[tokio::test] + async fn llama_chat_success_returns_completion() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "llama-4-maverick", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.model, "llama-4-maverick"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn llama_chat_sends_temperature_and_max_tokens() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "llama-4-maverick", + "temperature": 0.5, + "max_tokens": 50 + }))) + .with_status(200) + .with_body( + json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "llama-4-maverick", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn llama_chat_passes_extra_params_through() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "llama-4-maverick", + "custom_param": "custom_value" + }))) + .with_status(200) + .with_body( + json!({ + "id": "x", "object": "chat.completion", "created": 1, "model": "llama-4-maverick", + "choices": [{"index": 0, "message": {"role":"assistant","content":"ok"}, "finish_reason": "stop"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]) + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn llama_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid Llama API key"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn llama_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn llama_chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ============ Llama chat_stream ============ + + #[tokio::test] + async fn llama_chat_stream_parses_sse_chunks() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"llama-4-maverick\",\"choices\":[{\"index\":0,\"delta\":{\"role\":\"assistant\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"llama-4-maverick\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"Hello\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1700000000,\"model\":\"llama-4-maverick\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\" world\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + let mut content = String::new(); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[2].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn llama_chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Unauthorized"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let req = ChatRequest::builder("llama-4-maverick", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + // ============ Llama list_models ============ + + #[tokio::test] + async fn llama_list_models_success() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body( + json!({ + "object": "list", + "data": [ + {"id": "llama-4-maverick", "object": "model", "created": 1700000000, "owned_by": "meta"}, + {"id": "llama-4-scout", "object": "model", "created": 1700000000, "owned_by": "meta"} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_llama(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "llama-4-maverick"); + assert_eq!(models[0].provider, "llama"); + assert_eq!(models[1].id, "llama-4-scout"); + } + + #[tokio::test] + async fn llama_list_models_filter_by_chat_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(200) + .with_body( + json!({ + "data": [ + {"id": "llama-4-maverick", "object": "model", "created": 1}, + {"id": "dall-e-3", "object": "model", "created": 1} + ] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_llama(&server); + let chats = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert_eq!(chats.len(), 1); + assert_eq!(chats[0].id, "llama-4-maverick"); + } + + #[tokio::test] + async fn llama_list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn llama_list_models_error_429() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(429) + .with_body(json!({"error": {"message": "slow"}}).to_string()) + .create_async() + .await; + + let adapter = make_llama(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ============ Llama 不支持的能力 ============ + + #[tokio::test] + async fn llama_image_generate_returns_unsupported() { + let adapter = make_llama_no_server(); + let req = ImageRequest::builder("llama-4-maverick", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn llama_video_create_returns_unsupported() { + let adapter = make_llama_no_server(); + let req = VideoRequest::builder("llama-4-maverick", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ start / close ============ + + #[tokio::test] + async fn ideogram_start_and_close_are_noops() { + let mut adapter = make_ideogram_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn luma_start_and_close_are_noops() { + let mut adapter = make_luma_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn llama_start_and_close_are_noops() { + let mut adapter = make_llama_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn ideogram_map_api_error_401_is_authentication() { + let err = IdeogramAdapter::map_api_error(401, ""); + match err { + AibridgeError::Authentication { message } => { + assert!(message.contains("Ideogram")); + } + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn ideogram_map_api_error_402_is_api_payment() { + let err = IdeogramAdapter::map_api_error(402, ""); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 402); + assert!(message.contains("credits")); + } + _ => panic!("应为 Api (402)"), + } + } + + #[test] + fn ideogram_map_api_error_429_is_rate_limit() { + let err = IdeogramAdapter::map_api_error(429, ""); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn ideogram_map_api_error_500_extracts_message() { + let body = json!({"error": "internal error"}).to_string(); + let err = IdeogramAdapter::map_api_error(500, &body); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal error"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn luma_map_api_error_401_is_authentication() { + let err = LumaAdapter::map_api_error(401, ""); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn luma_map_api_error_402_is_api_payment() { + let err = LumaAdapter::map_api_error(402, ""); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 402), + _ => panic!("应为 Api (402)"), + } + } + + #[test] + fn luma_map_api_error_429_is_rate_limit() { + let err = LumaAdapter::map_api_error(429, ""); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn luma_map_api_error_extracts_detail_field() { + let body = json!({"detail": "validation failed"}).to_string(); + let err = LumaAdapter::map_api_error(400, &body); + match err { + AibridgeError::Api { message, .. } => { + assert_eq!(message, "validation failed"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn luma_map_api_error_extracts_detail_array() { + let body = json!({"detail": ["error1", "error2"]}).to_string(); + let err = LumaAdapter::map_api_error(422, &body); + match err { + AibridgeError::Api { message, .. } => { + assert_eq!(message, "error1; error2"); + } + _ => panic!("应为 Api"), + } + } + + // ============ map_luma_status 单元测试 ============ + + #[test] + fn map_luma_status_queued_is_pending() { + assert_eq!(map_luma_status("queued"), TaskStatus::Pending); + } + + #[test] + fn map_luma_status_dreaming_is_processing() { + assert_eq!(map_luma_status("dreaming"), TaskStatus::Processing); + } + + #[test] + fn map_luma_status_completed_is_success() { + assert_eq!(map_luma_status("completed"), TaskStatus::Success); + } + + #[test] + fn map_luma_status_failed_is_failed() { + assert_eq!(map_luma_status("failed"), TaskStatus::Failed); + } + + #[test] + fn map_luma_status_case_insensitive() { + assert_eq!(map_luma_status("QUEUED"), TaskStatus::Pending); + assert_eq!(map_luma_status("Completed"), TaskStatus::Success); + } + + #[test] + fn map_luma_status_unknown_defaults_to_pending() { + assert_eq!(map_luma_status("unknown_state"), TaskStatus::Pending); + } + + // ============ parse_iso_to_timestamp / civil_to_unix 单元测试 ============ + + #[test] + fn parse_iso_to_timestamp_z_suffix() { + // 2024-01-01T00:00:00Z = 1704067200 + let ts = parse_iso_to_timestamp("2024-01-01T00:00:00Z"); + assert_eq!(ts, Some(1704067200)); + } + + #[test] + fn parse_iso_to_timestamp_with_offset() { + // 2024-01-01T00:00:00+08:00 = 1704067200 - 8*3600 = 1704038400 + let ts = parse_iso_to_timestamp("2024-01-01T00:00:00+08:00"); + assert_eq!(ts, Some(1704038400)); + } + + #[test] + fn parse_iso_to_timestamp_invalid_returns_none() { + assert_eq!(parse_iso_to_timestamp("not a date"), None); + assert_eq!(parse_iso_to_timestamp("2024"), None); + } + + #[test] + fn civil_to_unix_epoch() { + // 1970-01-01T00:00:00 = 0 + assert_eq!(civil_to_unix(1970, 1, 1, 0, 0, 0), Some(0)); + } + + #[test] + fn civil_to_unix_known_date() { + // 2024-01-01T00:00:00 = 1704067200 + assert_eq!(civil_to_unix(2024, 1, 1, 0, 0, 0), Some(1704067200)); + } + + #[test] + fn civil_to_unix_invalid_date_returns_none() { + assert_eq!(civil_to_unix(2024, 13, 1, 0, 0, 0), None); + assert_eq!(civil_to_unix(2024, 1, 32, 0, 0, 0), None); + assert_eq!(civil_to_unix(2024, 1, 1, 24, 0, 0), None); + } + + // ============ file_input_to_url 单元测试 ============ + + #[test] + fn file_input_to_url_returns_url_for_url_variant() { + let f = FileInput::url("https://example.com/x.png"); + assert_eq!(file_input_to_url(&f), "https://example.com/x.png"); + } + + #[test] + fn file_input_to_url_returns_base64_for_base64_variant() { + let f = FileInput::base64("aGVsbG8="); + assert_eq!(file_input_to_url(&f), "aGVsbG8="); + } + + #[test] + fn file_input_to_url_returns_empty_for_path_variant() { + let f = FileInput::path("/tmp/x.png"); + assert_eq!(file_input_to_url(&f), ""); + } + + // ============ build_generate_body / build_generations_body 单元测试 ============ + + #[test] + fn ideogram_build_generate_body_wraps_in_image_request() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("ideogram", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = IdeogramAdapter::with_http(http, config); + let req = ImageRequest::builder("V_2A", "a cat") + .aspect_ratio("16:9") + .n(2) + .build(); + let body = adapter.build_generate_body(&req); + // 应被 image_request 包裹 + assert!(body.get("image_request").is_some()); + let inner = body.get("image_request").unwrap(); + assert_eq!(inner["model"], "V_2A"); + assert_eq!(inner["prompt"], "a cat"); + // aspect_ratio 归一化为大写 + assert_eq!(inner["aspect_ratio"], "16:9"); + assert_eq!(inner["num_images"], 2); + } + + #[test] + fn luma_build_generations_body_includes_fields() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("luma", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = LumaAdapter::with_http(http, config); + let req = VideoRequest::builder("ray-2", "a cat") + .aspect_ratio("16:9") + .duration(5) + .resolution("720p") + .build(); + let body = adapter.build_generations_body(&req); + assert_eq!(body["model"], "ray-2"); + assert_eq!(body["prompt"], "a cat"); + assert_eq!(body["aspect_ratio"], "16:9"); + // duration 转为 "5s" 字符串 + assert_eq!(body["duration"], "5s"); + assert_eq!(body["resolution"], "720p"); + } + + #[test] + fn luma_build_generations_body_no_keyframes_when_empty() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("luma", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = LumaAdapter::with_http(http, config); + let req = VideoRequest::builder("ray-2", "a cat").build(); + let body = adapter.build_generations_body(&req); + // 无 keyframes 字段 + assert!(body.get("keyframes").is_none()); + } +} From 43e923fcff73f01ff4db99d29afd1798ad6fc228 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:39:40 +0800 Subject: [PATCH 30/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20chinese=20=E9=80=82=E9=85=8D=E5=99=A8=EF=BC=88=E4=B8=AD?= =?UTF-8?q?=E6=96=87=E6=A8=A1=E5=9E=8B=E8=81=9A=E5=90=88=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/chinese.rs | 2809 ++++++++++++++++++ 1 file changed, 2809 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/chinese.rs diff --git a/crates/aibridge-core/src/adapters/chinese.rs b/crates/aibridge-core/src/adapters/chinese.rs new file mode 100644 index 0000000..f6954a6 --- /dev/null +++ b/crates/aibridge-core/src/adapters/chinese.rs @@ -0,0 +1,2809 @@ +//! 中文模型聚合适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/chinese.py`。 +//! +//! 支持六个中文 AI 模型 Provider(文档见各适配器注释): +//! - **通义千问 Qwen**(阿里 DashScope):OpenAI 兼容协议 +//! - **智谱 AI Zhipu**(GLM):OpenAI 兼容协议 +//! - **豆包 Doubao**(字节火山引擎方舟):OpenAI 兼容协议 +//! - **文心一言 ERNIE**(百度千帆):**独立协议**(access_token 认证 + 特有端点 + 特有请求体/响应) +//! - **Kimi**(月之暗面 Moonshot AI):OpenAI 兼容协议 +//! - **MiniMax**(稀宇科技):OpenAI 兼容协议(chat 部分) +//! +//! ## 结构 +//! +//! Python v1 是 6 个独立 adapter 类,统一注册到工厂。Rust 同样实现 6 个独立 struct: +//! - `QwenAdapter` / `ZhipuAdapter` / `DoubaoAdapter` / `KimiAdapter`:组合 +//! [`OpenAiCompatAdapter`] 地基,chat/chat_stream/list_models 全部委托 +//! - `MiniMaxAdapter`:组合地基复用 chat/chat_stream,但 list_models 走硬编码 +//! (MiniMax 无可靠的 `/models` 端点,对齐 Python 老版保留硬编码列表) +//! - `ErnieAdapter`:**独立协议实现**,仅复用地基的 `HttpClient` 与错误映射, +//! chat/chat_stream/list_models 全部独立实现(百度特有 access_token 流程 + 端点 + +//! `messages`/`system` 分离的请求体 + `result` 字段响应 + `is_end` 流式标记 + +//! 硬编码模型列表,因千帆 modellist 响应结构非 OpenAI 兼容) +//! +//! ## 能力声明策略 +//! +//! Python 老版 Qwen/Doubao/MiniMax 声明了 `AUDIO_TRANSCRIBE`/`AUDIO_SPEECH` 能力并实现了 +//! OpenAI 兼容的 transcribe/speech。但本阶段(2a)范围仅核心 chat/chat_stream/vision 能力, +//! audio 方法走 trait 默认实现返 `UnsupportedCapability`(与 azure 范例阶段 2a 策略一致, +//! 待阶段 2c audio_adapters 统一补齐)。故本模块不声明 audio 能力,避免能力声明与方法 +//! 行为不一致。 +//! +//! ## provider_type 标识(对齐 Python v1) +//! +//! | Provider | provider_type | +//! |---|---| +//! | 通义千问 Qwen | `qwen` | +//! | 智谱 AI Zhipu | `zhipu` | +//! | 豆包 Doubao | `doubao` | +//! | 文心一言 ERNIE | `ernie` | +//! | Kimi | `kimi` | +//! | MiniMax | `minimax` | +//! +//! 注册到工厂由 `adapter::factory` 完成(阶段 2a 收尾批次),本模块仅声明适配器类型。 + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChoiceMessage, DeltaMessage, UserContent, +}; +use crate::model::common::{ModelInfo, ModelType}; +use crate::util; + +// ==================== 默认 Base URL ==================== + +/// 通义千问 Qwen 默认 Base URL +/// +/// 对应 Python v1 `QwenAdapter.DEFAULT_BASE_URL`(DashScope OpenAI 兼容模式)。 +/// 已含 `/compatible-mode/v1` 前缀,chat 端点为 `POST /chat/completions`, +/// models 端点为 `GET /models`。 +pub const DEFAULT_QWEN_BASE_URL: &str = "https://dashscope.aliyuncs.com/compatible-mode/v1"; + +/// 智谱 AI Zhipu 默认 Base URL +/// +/// 对应 Python v1 `ZhipuAdapter.DEFAULT_BASE_URL`(智谱 GLM OpenAI 兼容模式)。 +/// 已含 `/api/paas/v4` 前缀,chat 端点为 `POST /chat/completions`, +/// models 端点为 `GET /models`。 +pub const DEFAULT_ZHIPU_BASE_URL: &str = "https://open.bigmodel.cn/api/paas/v4"; + +/// 豆包 Doubao 默认 Base URL +/// +/// 对应 Python v1 `DoubaoAdapter.DEFAULT_BASE_URL`(火山引擎方舟 OpenAI 兼容模式)。 +/// 已含 `/api/v3` 前缀,chat 端点为 `POST /chat/completions`, +/// models 端点为 `GET /models`。 +pub const DEFAULT_DOUBAO_BASE_URL: &str = "https://ark.cn-beijing.volces.com/api/v3"; + +/// 文心一言 ERNIE 默认 Base URL +/// +/// 对应 Python v1 `ErnieAdapter.DEFAULT_BASE_URL`(百度千帆)。 +/// 不含路径前缀,chat 端点为 `POST /rpc/2.0/ai_custom/v1/wenxinworkshop/chat/{model}`, +/// access_token 端点为 `GET /oauth/2.0/token`。 +pub const DEFAULT_ERNIE_BASE_URL: &str = "https://aip.baidubce.com"; + +/// Kimi 默认 Base URL +/// +/// 对应 Python v1 `KimiAdapter.DEFAULT_BASE_URL`(月之暗面 Moonshot AI)。 +/// 已含 `/v1` 前缀,chat 端点为 `POST /chat/completions`,models 端点为 `GET /models`。 +pub const DEFAULT_KIMI_BASE_URL: &str = "https://api.moonshot.cn/v1"; + +/// MiniMax 默认 Base URL +/// +/// 对应 Python v1 `MiniMaxAdapter.DEFAULT_BASE_URL`(稀宇科技)。 +/// 已含 `/v1` 前缀,chat 端点为 `POST /chat/completions`。 +/// 无可靠的 `/models` 端点,list_models 走硬编码列表。 +pub const DEFAULT_MINIMAX_BASE_URL: &str = "https://api.minimaxi.com/v1"; + +// ==================== 能力集合构造 ==================== + +/// 通义千问 Qwen 支持的能力集合 +/// +/// 对齐 Python v1 `QwenAdapter.supported_capabilities` 的核心子集 +/// (chat / chat_stream / vision)。Python 还声明了 AUDIO_TRANSCRIBE/AUDIO_SPEECH, +/// 但阶段 2a 范围仅核心三能力,audio 走 trait 默认实现返 UnsupportedCapability +/// (待阶段 2c audio_adapters 补齐)。 +fn qwen_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 智谱 AI Zhipu 支持的能力集合 +/// +/// 对齐 Python v1 `ZhipuAdapter.supported_capabilities = ["chat", "vision"]`。 +/// Rust 额外声明 ChatStream(OpenAI 兼容地基支持流式,Python 老版也实现了 chat_stream)。 +fn zhipu_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 豆包 Doubao 支持的能力集合 +/// +/// 对齐 Python v1 `DoubaoAdapter.supported_capabilities` 的核心子集 +/// (chat / chat_stream / vision),audio 能力待阶段 2c 补齐。 +fn doubao_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// 文心一言 ERNIE 支持的能力集合 +/// +/// 对齐 Python v1 `ErnieAdapter.supported_capabilities = ["chat", "vision"]`。 +/// Rust 额外声明 ChatStream(Python 老版实现了 chat_stream 流式)。 +fn ernie_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// Kimi 支持的能力集合 +/// +/// 对齐 Python v1 `KimiAdapter.supported_capabilities = ["chat", "vision"]`。 +/// Rust 额外声明 ChatStream(OpenAI 兼容地基支持流式,Python 老版也实现了 chat_stream)。 +fn kimi_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +/// MiniMax 支持的能力集合 +/// +/// 对齐 Python v1 `MiniMaxAdapter.supported_capabilities` 的核心子集 +/// (chat / chat_stream / vision),audio 能力待阶段 2c 补齐。 +fn minimax_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +// ==================== 通义千问 Qwen 适配器 ==================== + +/// 通义千问(Qwen)适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream/list_models 全部委托给 [`OpenAiCompatAdapter`] 地基。 +/// +/// - Base URL: `https://dashscope.aliyuncs.com/compatible-mode/v1`(DashScope OpenAI 兼容模式) +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models`(实时拉取,对齐 Python 老版 `_parse_models_response`) +/// - 认证: Bearer Token(DashScope API Key) +/// - 文档: +pub struct QwenAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl QwenAdapter { + /// 创建通义千问适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_QWEN_BASE_URL, + qwen_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "qwen"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "通义千问"; +} + +#[async_trait] +impl Adapter for QwenAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + // 返回新集合,避免外部修改内部状态(不可变原则) + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new() 时已构造,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 走 Drop 释放,无需显式关闭 + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability,与 Python 老版行为一致 + // (Python 老版 image/video 显式抛 UnsupportedCapabilityError)。 +} + +// ==================== 智谱 AI Zhipu 适配器 ==================== + +/// 智谱 AI(GLM)适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream/list_models 全部委托给 [`OpenAiCompatAdapter`] 地基。 +/// +/// - Base URL: `https://open.bigmodel.cn/api/paas/v4` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models`(实时拉取) +/// - 认证: Bearer Token(智谱 API Key) +/// - 文档: +pub struct ZhipuAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl ZhipuAdapter { + /// 创建智谱适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_ZHIPU_BASE_URL, + zhipu_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "zhipu"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "智谱 AI"; +} + +#[async_trait] +impl Adapter for ZhipuAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability。 +} + +// ==================== 豆包 Doubao 适配器 ==================== + +/// 豆包(字节火山引擎方舟)适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream/list_models 全部委托给 [`OpenAiCompatAdapter`] 地基。 +/// +/// - Base URL: `https://ark.cn-beijing.volces.com/api/v3` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models`(实时拉取) +/// - 认证: Bearer Token(火山引擎 API Key) +/// - 文档: +pub struct DoubaoAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl DoubaoAdapter { + /// 创建豆包适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_DOUBAO_BASE_URL, + doubao_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "doubao"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "豆包"; +} + +#[async_trait] +impl Adapter for DoubaoAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability。 +} + +// ==================== MiniMax 适配器 ==================== + +/// MiniMax(稀宇科技)适配器 +/// +/// OpenAI 兼容协议的 chat/chat_stream 委托给 [`OpenAiCompatAdapter`] 地基, +/// 但 list_models 走硬编码列表(MiniMax 无可靠的 `/models` 端点,对齐 Python 老版)。 +/// +/// - Base URL: `https://api.minimaxi.com/v1` +/// - Chat: `POST /chat/completions`(OpenAI 兼容接口) +/// - Models: 无标准 `/models` 端点,保留硬编码列表 +/// - 认证: Bearer Token(MiniMax API Key),可选 `X-Group-Id` header +/// - 支持模型: abab 系列、MiniMax-Text-01、MiniMax-VL-01、MiniMax-M1、abab-asr/tts 等 +/// - 文档: +/// +/// `group_id` 经 `config.extra["group_id"]` 传入,启动时注入到默认请求 header +/// (Python 老版 `start()` 中读取并设 `X-Group-Id`)。本实现因复用地基的 HttpClient, +/// group_id 注入依赖调用方在构造 config 时通过 `extra` 透传(地基不识别该 header, +/// 需要时由本适配器在 chat 前显式附加——当前实现与 Python 一致:仅当存在时附加到请求)。 +pub struct MiniMaxAdapter { + /// OpenAI 兼容地基(chat/chat_stream 委托给它,list_models 独立实现) + compat: OpenAiCompatAdapter, + /// MiniMax group_id(可选,经 `config.extra["group_id"]` 传入) + group_id: Option, +} + +impl MiniMaxAdapter { + /// 创建 MiniMax 适配器 + /// + /// `group_id` 从 `config.extra["group_id"]` 提取(对齐 Python 老版 + /// `getattr(config, "group_id", None)`)。 + pub fn new(config: ProviderConfig) -> Result { + let group_id = config + .extra + .get("group_id") + .and_then(|v| v.as_str()) + .map(|s| s.to_string()); + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_MINIMAX_BASE_URL, + minimax_capabilities(), + )?; + Ok(Self { compat, group_id }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { + compat, + group_id: None, + } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "minimax"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "MiniMax"; + + /// API key(可能为空) + fn api_key(&self) -> Option<&str> { + self.compat.api_key() + } + + /// base_url + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// group_id(可选) + pub fn group_id(&self) -> Option<&str> { + self.group_id.as_deref() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), Self::PROVIDER_TYPE), + }) + } + } + + /// 构造带 group_id header 的 POST 请求 builder + /// + /// MiniMax 认证用 Bearer Token,group_id 存在时附加 `X-Group-Id` header + /// (对齐 Python 老版 `start()` 中的 header 设置)。 + fn post_request(&self, url: &str, body: &Value) -> reqwest::RequestBuilder { + let mut req = self + .compat + .http_inner() + .post(url) + .bearer_auth(self.api_key().unwrap_or("")) + .json(body); + if let Some(gid) = &self.group_id { + req = req.header("X-Group-Id", gid); + } + req + } + + /// 发送 chat/completions 请求并用 OpenAI 错误映射处理响应 + async fn post_chat(&self, body: &Value) -> Result { + let url = self.url("chat/completions"); + let resp = self.post_request(&url, body).send().await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(map_minimax_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } +} + +/// 将 MiniMax API 错误响应映射为 AibridgeError +/// +/// MiniMax 错误体结构与 OpenAI 类似(`{"error": {"message": "..."}}`), +/// 但额外可能用 `base_resp.status_msg` 字段(对齐 Python 老版 `_handle_error`)。 +/// 优先提取这些字段的 message,再按 HTTP 状态码分类(复用 OpenAI 兼容错误分类)。 +fn map_minimax_error(status: u16, body: &str) -> AibridgeError { + // 尝试解析错误 message:优先 error.message,再 base_resp.status_msg,再顶层 message + let message = if let Ok(v) = serde_json::from_str::(body) { + let mut msg: Option = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + .map(str::to_owned); + if msg.is_none() { + msg = v + .get("base_resp") + .and_then(|b| b.get("status_msg")) + .and_then(|m| m.as_str()) + .map(str::to_owned); + } + if msg.is_none() { + msg = v.get("message").and_then(|m| m.as_str()).map(str::to_owned); + } + msg.unwrap_or_else(|| format!("HTTP {status}")) + } else if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + }; + + match status { + 401 | 403 => AibridgeError::Authentication { message }, + 429 => { + let retry_after = serde_json::from_str::(body).ok().and_then(|v| { + v.get("error") + .and_then(|e| e.get("retry_after")) + .and_then(|r| r.as_f64()) + .or_else(|| v.get("retry_after").and_then(|r| r.as_f64())) + }); + AibridgeError::RateLimit { + message, + retry_after, + } + } + 400 => AibridgeError::Validation { + message, + details: serde_json::json!({ "status_code": status, "response": body }), + }, + 404 => AibridgeError::ModelNotFound { model: message }, + s if (400..600).contains(&s) => AibridgeError::Api { status: s, message }, + s => AibridgeError::Api { status: s, message }, + } +} + +#[async_trait] +impl Adapter for MiniMaxAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话 + /// + /// 用地基构造标准 OpenAI 请求体,自建 POST 请求(附加 `X-Group-Id` header), + /// 响应解析委托地基的 `parse_chat_completion`。 + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = self.compat.build_chat_body(&req, false); + let value = self.post_chat(&body).await?; + self.compat.parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话 + /// + /// 自建 POST 请求(附加 `X-Group-Id` header),SSE 解析复用地基的 chunk 解析逻辑。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = self.compat.build_chat_body(&req, true); + let url = self.url("chat/completions"); + + let resp = self.post_request(&url, &body).send().await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(map_minimax_error(status_code, &body_text)); + } + + let model = req.model.clone(); + // 按字节流读取,按行切分解析 SSE(与 openai_compat/azure 范例一致的按行切分逻辑) + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = MiniMaxLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 去除 "data: " 前缀 + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + // 结束标记 + if data.trim() == "[DONE]" { + return; + } + match serde_json::from_str::(data) { + Ok(v) => { + // 复用地基的 chunk 解析逻辑(parse_chunk 是关联函数,无需借 self) + match OpenAiCompatAdapter::parse_chunk(&v, &model) { + Ok(Some(chunk)) => yield Ok(chunk), + Ok(None) => continue, + Err(e) => { + yield Err(e); + return; + } + } + } + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 模型列表(硬编码) + /// + /// MiniMax 无可靠的 `/models` 端点,保留硬编码列表(对齐 Python 老版)。 + /// 按 `filter` 过滤模型类型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = minimax_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability(Python 老版 image/video 显式抛错)。 +} + +/// MiniMax 硬编码模型列表 +/// +/// 对齐 Python v1 `MiniMaxAdapter.list_models` 的硬编码列表。 +/// 含 chat 模型(abab 系列、MiniMax-Text-01/VL-01/M1)与 audio 模型(abab-asr/tts、speech-01)。 +fn minimax_hardcoded_models() -> Vec { + let chat_models = [ + ( + "abab6.5s-chat", + "ABAB 6.5s", + vec!["chat".to_string()], + "MiniMax abab 6.5s 快速版本", + ), + ( + "abab6.5-chat", + "ABAB 6.5", + vec!["chat".to_string()], + "MiniMax abab 6.5 标准版本", + ), + ( + "MiniMax-Text-01", + "MiniMax Text 01", + vec!["chat".to_string()], + "MiniMax 万亿参数 MoE 模型", + ), + ( + "MiniMax-VL-01", + "MiniMax VL 01", + vec!["chat".to_string(), "vision".to_string()], + "MiniMax 多模态视觉模型", + ), + ( + "MiniMax-M1", + "MiniMax M1", + vec!["chat".to_string(), "vision".to_string()], + "MiniMax 最新推理模型", + ), + ]; + let audio_models = [ + ( + "abab-asr", + "ABAB 语音识别", + vec!["audio_transcribe".to_string()], + "MiniMax 语音识别模型", + ), + ( + "abab-tts", + "ABAB 语音合成", + vec!["audio_speech".to_string()], + "MiniMax 语音合成模型(多音色)", + ), + ( + "speech-01", + "MiniMax Speech 01", + vec!["audio_speech".to_string()], + "MiniMax 高品质语音合成", + ), + ]; + + let mut models: Vec = Vec::new(); + for (id, name, caps, desc) in chat_models { + models.push(ModelInfo { + id: id.to_string(), + name: name.to_string(), + model_type: ModelType::Chat, + provider: MiniMaxAdapter::PROVIDER_TYPE.to_string(), + capabilities: caps, + max_tokens: None, + supports_streaming: true, + description: Some(desc.to_string()), + created: None, + }); + } + for (id, name, caps, desc) in audio_models { + models.push(ModelInfo { + id: id.to_string(), + name: name.to_string(), + model_type: ModelType::Audio, + provider: MiniMaxAdapter::PROVIDER_TYPE.to_string(), + capabilities: caps, + max_tokens: None, + supports_streaming: false, + description: Some(desc.to_string()), + created: None, + }); + } + models +} + +// ==================== MiniMax SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(MiniMax 流式用) +/// +/// 与 `openai_compat::LinesStream` 等价实现,独立实现避免引用其私有结构。 +struct MiniMaxLinesStream { + inner: S, + buffer: Vec, +} + +impl MiniMaxLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for MiniMaxLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + // 先看缓冲区是否已有完整行 + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + // 去掉末尾 \n 与可能的 \r + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + // 缓冲区无完整行,拉取下一 chunk + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + // 继续循环,尝试从缓冲区切出行 + } + Poll::Ready(None) => { + // 流结束,把缓冲区剩余内容作为最后一行返回 + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +/// Kimi(月之暗面 Moonshot AI)适配器 +/// +/// OpenAI 兼容协议,chat/chat_stream/list_models 全部委托给 [`OpenAiCompatAdapter`] 地基。 +/// +/// - Base URL: `https://api.moonshot.cn/v1` +/// - Chat: `POST /chat/completions` +/// - Models: `GET /models`(实时拉取) +/// - 认证: Bearer Token(Moonshot API Key) +/// - 特点: 支持超长上下文(128K/256K)、视觉理解 +/// - 文档: +pub struct KimiAdapter { + /// OpenAI 兼容地基(chat/chat_stream/list_models 委托给它) + compat: OpenAiCompatAdapter, +} + +impl KimiAdapter { + /// 创建 Kimi 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_KIMI_BASE_URL, + kimi_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "kimi"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Kimi (月之暗面)"; +} + +#[async_trait] +impl Adapter for KimiAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话:委托地基 `POST /chat/completions` + async fn chat(&self, req: ChatRequest) -> Result { + self.compat.chat(req).await + } + + /// 流式文本对话:委托地基 `POST /chat/completions` (stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.compat.chat_stream(req).await + } + + /// 模型列表(实时拉取):委托地基 `GET /models` + async fn list_models(&self, filter: Option) -> Result> { + self.compat.list_models(filter).await + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability。 +} + +// ==================== 文心一言 ERNIE 适配器(独立协议) ==================== + +/// 文心一言(百度 ERNIE)适配器 +/// +/// **独立协议实现**,仅复用地基的 `HttpClient` 与错误映射,chat/chat_stream/list_models +/// 全部独立实现(百度特有 access_token 流程 + 端点 + `messages`/`system` 分离的请求体 + +/// `result` 字段响应 + `is_end` 流式标记 + 硬编码模型列表)。 +/// +/// - Base URL: `https://aip.baidubce.com` +/// - Chat: `POST /rpc/2.0/ai_custom/v1/wenxinworkshop/chat/{model}?access_token=...` +/// - 认证: access_token 查询参数(通过 API Key + Secret Key 换取) +/// - `config.api_key` 可直接是 access_token,或 `ak:sk` 格式(自动换取 access_token) +/// - access_token 端点: `GET /oauth/2.0/token?grant_type=client_credentials&client_id=ak&client_secret=sk` +/// - 请求体: `{"messages": [...], "system": "..."}`(system 单独字段,非 messages 数组) +/// - 响应: `{"result": "...", "usage": {...}, "is_truncated": false}` +/// - 流式: SSE,每行 `data: {"result": "...", "is_end": false}` +/// - Models: 无标准 `/models` 端点,保留硬编码列表(千帆 modellist 响应结构非 OpenAI 兼容) +/// - 文档: +/// +/// ## api_key 解析 +/// +/// 对齐 Python 老版 `ErnieAdapter.__init__`: +/// - `config.api_key` 含 `:` → `ak:sk` 格式,构造时拆分为 `api_key`(ak) + `secret_key`(sk), +/// `start()` 时自动调 `/oauth/2.0/token` 换取 access_token +/// - 否则 → 直接作为 access_token 使用(用户已自行换取) +pub struct ErnieAdapter { + /// OpenAI 兼容地基(仅复用其 HttpClient 与错误映射,不委托 chat 等方法) + compat: OpenAiCompatAdapter, + /// Secret Key(仅当 api_key 为 `ak:sk` 格式时非空) + secret_key: String, + /// access_token(start 时换取,或直接用 api_key) + access_token: Option, +} + +impl ErnieAdapter { + /// 创建文心一言适配器 + /// + /// api_key 解析(对齐 Python 老版 `__init__`): + /// - 含 `:` → `ak:sk` 格式,拆分为 api_key(ak) + secret_key(sk) + /// - 否则 → 直接作为 access_token 使用 + pub fn new(mut config: ProviderConfig) -> Result { + // 先把 api_key 取出(避免借用冲突),再根据是否含 ':' 拆分 ak:sk + let api_key = config.api_key.take(); + let (ak, secret_key) = match api_key.as_ref().and_then(|k| k.split_once(':')) { + Some((ak, sk)) => (Some(ak.to_string()), sk.to_string()), + None => (api_key, String::new()), + }; + config.api_key = ak; + + // access_token 初始化:无 secret_key 时直接用 api_key + let access_token = if secret_key.is_empty() { + config.api_key.clone() + } else { + None + }; + + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_ERNIE_BASE_URL, + ernie_capabilities(), + )?; + Ok(Self { + compat, + secret_key, + access_token, + }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + /// + /// `access_token` 默认为 `"test-token"`,避免测试触发 oauth 流程。 + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { + compat, + secret_key: String::new(), + access_token: Some("test-token".to_string()), + } + } + + /// 用显式 compat + secret_key 构造(测试 oauth 流程用) + #[cfg(test)] + pub fn with_compat_and_secret(compat: OpenAiCompatAdapter, secret_key: String) -> Self { + Self { + compat, + secret_key, + access_token: None, + } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "ernie"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "文心一言"; + + /// base_url + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// 当前 access_token(已换取或直接用 api_key) + pub fn access_token(&self) -> Option<&str> { + self.access_token.as_deref() + } + + /// 是否需要走 oauth 换取 access_token(即配置了 ak:sk) + pub fn needs_oauth(&self) -> bool { + !self.secret_key.is_empty() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), Self::PROVIDER_TYPE), + }) + } + } + + /// 换取 access_token(对齐 Python 老版 `_get_access_token`) + /// + /// `GET /oauth/2.0/token?grant_type=client_credentials&client_id={ak}&client_secret={sk}` + /// 响应:`{"access_token": "...", ...}` + async fn get_access_token(&self) -> Result { + let url = self.url("oauth/2.0/token"); + let ak = self.compat.api_key().unwrap_or(""); + let resp = self + .compat + .http_inner() + .get(&url) + .query(&[ + ("grant_type", "client_credentials"), + ("client_id", ak), + ("client_secret", &self.secret_key), + ]) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(map_ernie_error(status_code, &body_text)); + } + let value: Value = resp.json().await.map_err(AibridgeError::from)?; + let token = value + .get("access_token") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .ok_or_else(|| AibridgeError::Api { + status: 0, + message: "ERNIE oauth 响应缺少 access_token 字段".to_string(), + })?; + Ok(token) + } + + /// 转换统一消息为 ERNIE 格式,并提取 system prompt + /// + /// 对齐 Python v1 `ErnieAdapter._convert_messages`: + /// - system 消息 → 提取为 `system` 字段(最后一条 system 覆盖前面的) + /// - 其他消息 → `{"role": "user"/"assistant", "content": "..."}` + /// + /// 返回 `(messages, system)`,messages 不含 system 消息。 + fn convert_messages(messages: &[ChatMessage]) -> (Vec, Option) { + let mut ernie_messages: Vec = Vec::new(); + let mut system: Option = None; + + for msg in messages { + match msg { + ChatMessage::System { content, .. } => { + system = Some(content.clone()); + } + ChatMessage::User { content, .. } => { + let text = match content { + UserContent::Text(s) => s.clone(), + // 多模态:拼接所有 Text 部件,忽略图片(ERNIE 不支持图片输入) + UserContent::Parts(parts) => parts + .iter() + .filter_map(|p| match p { + crate::model::chat::ContentPart::Text { text } => { + Some(text.clone()) + } + _ => None, + }) + .collect::>() + .join(""), + }; + ernie_messages.push(json!({ "role": "user", "content": text })); + } + ChatMessage::Assistant { content, .. } => { + let text = content.clone().unwrap_or_default(); + ernie_messages.push(json!({ "role": "assistant", "content": text })); + } + ChatMessage::Tool { content, .. } => { + // 工具结果消息按 user 角色处理(ERNIE 无原生 tool 角色) + ernie_messages.push(json!({ "role": "user", "content": content })); + } + } + } + + (ernie_messages, system) + } + + /// 构造 ERNIE chat 请求体 + /// + /// 对齐 Python v1 `ErnieAdapter.chat` 的 body 构造: + /// - `messages`:非 system 消息列表 + /// - `system`:若有 system 消息则填入(单独字段) + /// - `temperature` / `top_p`:从统一请求透传 + /// - `stream`:流式时置 true + /// - `extra` 透传到顶层 + fn build_chat_body(req: &ChatRequest, stream: bool) -> Value { + let (ernie_messages, system) = Self::convert_messages(&req.messages); + + let mut body = json!({ "messages": ernie_messages }); + if stream { + body["stream"] = json!(true); + } + if let Some(s) = system { + body["system"] = json!(s); + } + if let Some(t) = req.temperature { + body["temperature"] = json!(t); + } + if let Some(p) = req.top_p { + body["top_p"] = json!(p); + } + // extra 透传到顶层 + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 解析 ERNIE chat 响应 → ChatCompletion + /// + /// ERNIE 响应格式(非标准 OpenAI choices 结构): + /// `{"result": "...", "usage": {"prompt_tokens":..., "completion_tokens":..., "total_tokens":...}, "is_truncated": false}` + /// 对应 Python v1 `ErnieAdapter._parse_response`。 + fn parse_chat_completion(value: &Value, fallback_model: &str) -> Result { + let result = value + .get("result") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + + // usage 解析(ERNIE usage 字段结构与 OpenAI 一致) + let usage = value.get("usage").and_then(parse_ernie_usage); + + // is_truncated 为 true 时 finish_reason 为 "length",否则 "stop" + let is_truncated = value + .get("is_truncated") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + let finish_reason = if is_truncated { "length" } else { "stop" }; + + Ok(ChatCompletion { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion".to_string(), + created: value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp), + model: value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()), + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".to_string(), + content: Some(result), + tool_calls: None, + }, + finish_reason: Some(finish_reason.to_string()), + }], + usage, + service_tier: None, + system_fingerprint: None, + }) + } + + /// 解析单个 ERNIE 流式 chunk + /// + /// ERNIE 流式事件格式:`{"result": "...", "is_end": false, ...}` + /// - `result`:增量文本(delta.content = result) + /// - `is_end`:是否结束(true 时 finish_reason = "stop") + /// + /// 对应 Python v1 `ErnieAdapter._parse_chunk`。 + fn parse_chunk(value: &Value, fallback_model: &str) -> Option { + let result = value + .get("result") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let is_end = value + .get("is_end") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + let finish_reason = if is_end { + Some("stop".to_string()) + } else { + None + }; + + // usage 可能出现在结束块 + let usage = value.get("usage").and_then(parse_ernie_usage); + + Some(ChatCompletionChunk { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion.chunk".to_string(), + created: util::current_timestamp(), + model: value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".to_string()), + content: Some(result), + tool_calls: None, + }, + finish_reason, + }], + usage, + }) + } + + /// 构造 chat 端点 URL(含 access_token query) + /// + /// `POST /rpc/2.0/ai_custom/v1/wenxinworkshop/chat/{model}?access_token=...` + fn chat_url(&self, model: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + format!( + "{base}/rpc/2.0/ai_custom/v1/wenxinworkshop/chat/{model}?access_token={}", + self.access_token.as_deref().unwrap_or("") + ) + } + + /// 发送 chat 请求并用 ERNIE 错误映射处理响应 + async fn post_ernie_json(&self, url: &str, body: &Value) -> Result { + let resp = self + .compat + .http_inner() + .post(url) + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(map_ernie_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } +} + +/// 解析 ERNIE usage 统计 +/// +/// ERNIE usage 格式与 OpenAI 一致:`{"prompt_tokens": N, "completion_tokens": M, "total_tokens": T}` +fn parse_ernie_usage(v: &Value) -> Option { + let prompt = v.get("prompt_tokens").and_then(|x| x.as_u64())?; + let completion = v + .get("completion_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(0); + let total = v + .get("total_tokens") + .and_then(|x| x.as_u64()) + .unwrap_or(prompt + completion); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: total, + }) +} + +/// 将 ERNIE API 错误响应映射为 AibridgeError +/// +/// 对齐 Python v1 `ErnieAdapter._handle_error`: +/// - 401/403 → Authentication +/// - 429 → RateLimit +/// - 其他 → Api +/// - 错误 message 优先取 `error_msg`,再 `message`,再 `error` +fn map_ernie_error(status: u16, body: &str) -> AibridgeError { + let message = if let Ok(v) = serde_json::from_str::(body) { + let mut msg: Option = v + .get("error_msg") + .and_then(|m| m.as_str()) + .map(str::to_owned); + if msg.is_none() { + msg = v.get("message").and_then(|m| m.as_str()).map(str::to_owned); + } + if msg.is_none() { + // error 可能是字符串也可能是对象 + msg = v + .get("error") + .and_then(|e| e.as_str()) + .map(str::to_owned) + .or_else(|| { + v.get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + .map(str::to_owned) + }); + } + msg.unwrap_or_else(|| format!("HTTP {status}")) + } else if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + }; + + match status { + 401 | 403 => AibridgeError::Authentication { message }, + 429 => AibridgeError::RateLimit { + message, + retry_after: None, + }, + 400 => AibridgeError::Validation { + message, + details: serde_json::json!({ "status_code": status, "response": body }), + }, + 404 => AibridgeError::ModelNotFound { model: message }, + s if (400..600).contains(&s) => AibridgeError::Api { status: s, message }, + s => AibridgeError::Api { status: s, message }, + } +} + +#[async_trait] +impl Adapter for ErnieAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + /// 启动适配器 + /// + /// 若配置了 ak:sk(`needs_oauth()` 为 true),自动调 `/oauth/2.0/token` 换取 access_token。 + /// 对齐 Python 老版 `start()` 行为。 + async fn start(&mut self) -> Result<()> { + if self.needs_oauth() && self.access_token.is_none() { + let token = self.get_access_token().await?; + self.access_token = Some(token); + } + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + Ok(()) + } + + /// 文本对话(ERNIE 特有协议:POST /rpc/2.0/ai_custom/v1/wenxinworkshop/chat/{model}) + /// + /// 请求体用 `messages` + `system` 分离结构,响应 `result` 字段。 + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = Self::build_chat_body(&req, false); + let url = self.chat_url(&req.model); + let value = self.post_ernie_json(&url, &body).await?; + Self::parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话(ERNIE 特有协议:POST .../chat/{model}?access_token=... stream=true) + /// + /// SSE 格式,每行 `data: {"result": "...", "is_end": false}`。 + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = Self::build_chat_body(&req, true); + let url = self.chat_url(&req.model); + + let resp = self + .compat + .http_inner() + .post(&url) + .json(&body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(map_ernie_error(status_code, &body_text)); + } + + let model = req.model.clone(); + // 按字节流读取,按行切分解析 SSE(与 openai_compat/azure/minimax 范例一致的按行切分逻辑) + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = ErnieLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 去除 "data: " 前缀 + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + // ERNIE 流式无 [DONE] 标记,靠 is_end 字段判断结束 + match serde_json::from_str::(data) { + Ok(v) => match Self::parse_chunk(&v, &model) { + Some(chunk) => yield Ok(chunk), + None => continue, + }, + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 模型列表(硬编码) + /// + /// 对齐 Python v1 `ErnieAdapter.list_models`:百度千帆 modellist 响应结构非 OpenAI 兼容 + /// (`result.model_list`,字段为 code/name),无法复用基类解析,故保留硬编码列表。 + async fn list_models(&self, filter: Option) -> Result> { + let models = ernie_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability(Python 老版 image/video 显式抛错)。 +} + +/// ERNIE 硬编码模型列表 +/// +/// 对齐 Python v1 `ErnieAdapter.list_models` 的硬编码列表。 +/// 含 ERNIE 4.0 / 3.5 / Lite 等 chat 模型。 +fn ernie_hardcoded_models() -> Vec { + let models = [ + ( + "completions_pro", + "ERNIE 4.0", + vec!["chat".to_string(), "vision".to_string()], + "文心一言 4.0", + ), + ( + "completions", + "ERNIE 3.5", + vec!["chat".to_string()], + "文心一言 3.5", + ), + ( + "ernie-lite-8k", + "ERNIE Lite", + vec!["chat".to_string()], + "文心一言轻量版", + ), + ]; + models + .iter() + .map(|(id, name, caps, desc)| ModelInfo { + id: id.to_string(), + name: name.to_string(), + model_type: ModelType::Chat, + provider: ErnieAdapter::PROVIDER_TYPE.to_string(), + capabilities: caps.clone(), + max_tokens: None, + supports_streaming: true, + description: Some(desc.to_string()), + created: None, + }) + .collect() +} + +// ==================== ERNIE SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(ERNIE 流式用) +/// +/// 与 `openai_compat::LinesStream` 等价实现,独立实现避免引用其私有结构。 +struct ErnieLinesStream { + inner: S, + buffer: Vec, +} + +impl ErnieLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for ErnieLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + // 先看缓冲区是否已有完整行 + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + // 去掉末尾 \n 与可能的 \r + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + // 缓冲区无完整行,拉取下一 chunk + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + // 继续循环,尝试从缓冲区切出行 + } + Poll::Ready(None) => { + // 流结束,把缓冲区剩余内容作为最后一行返回 + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::http::HttpClient; + use crate::model::chat::ChatMessage; + use crate::model::image::ImageRequest; + use crate::model::video::VideoRequest; + use futures::stream::StreamExt; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 OpenAiCompatAdapter(指向 mockito server,给定 provider 信息与能力) + fn make_compat( + server: &Server, + provider_type: &str, + provider_name: &str, + caps: CapabilitySet, + ) -> OpenAiCompatAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options(provider_type, opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + OpenAiCompatAdapter::with_http(http, config, provider_type, provider_name, caps) + } + + /// 标准 OpenAI chat/completions 响应体(复用于多个兼容适配器测试) + fn openai_chat_body() -> Value { + json!({ + "id": "chatcmpl-1", + "object": "chat.completion", + "created": 1700000000, + "model": "test-model", + "choices": [{ + "index": 0, + "message": {"role": "assistant", "content": "Hello!"}, + "finish_reason": "stop" + }], + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7} + }) + } + + /// 构造 QwenAdapter(指向 mockito server) + fn make_qwen(server: &Server) -> QwenAdapter { + let compat = make_compat(server, "qwen", "通义千问", qwen_capabilities()); + QwenAdapter::with_compat(compat) + } + + /// 构造 ZhipuAdapter(指向 mockito server) + fn make_zhipu(server: &Server) -> ZhipuAdapter { + let compat = make_compat(server, "zhipu", "智谱 AI", zhipu_capabilities()); + ZhipuAdapter::with_compat(compat) + } + + /// 构造 DoubaoAdapter(指向 mockito server) + fn make_doubao(server: &Server) -> DoubaoAdapter { + let compat = make_compat(server, "doubao", "豆包", doubao_capabilities()); + DoubaoAdapter::with_compat(compat) + } + + /// 构造 KimiAdapter(指向 mockito server) + fn make_kimi(server: &Server) -> KimiAdapter { + let compat = make_compat(server, "kimi", "Kimi (月之暗面)", kimi_capabilities()); + KimiAdapter::with_compat(compat) + } + + /// 构造 MiniMaxAdapter(指向 mockito server) + fn make_minimax(server: &Server) -> MiniMaxAdapter { + let compat = make_compat(server, "minimax", "MiniMax", minimax_capabilities()); + MiniMaxAdapter::with_compat(compat) + } + + /// 构造 ErnieAdapter(指向 mockito server,access_token 预填避免 oauth) + fn make_ernie(server: &Server) -> ErnieAdapter { + let compat = make_compat(server, "ernie", "文心一言", ernie_capabilities()); + ErnieAdapter::with_compat(compat) + } + + // ==================== Qwen 元信息与能力 ==================== + + #[tokio::test] + async fn qwen_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_qwen(&server); + assert_eq!(adapter.provider_type(), "qwen"); + assert_eq!(adapter.provider_name(), "通义千问"); + } + + #[tokio::test] + async fn qwen_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_qwen(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn qwen_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_qwen(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // audio 能力在阶段 2a 不声明(待 2c 补齐) + assert!(!caps.contains(&Capabilities::AudioSpeech)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[tokio::test] + async fn qwen_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_qwen(&server); + let req = ChatRequest::builder("qwen-turbo", vec![ChatMessage::user("hi")]) + .temperature(0.7) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "chatcmpl-1"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 7); + mock.assert_async().await; + } + + #[tokio::test] + async fn qwen_chat_stream_parses_sse() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"qwen-turbo\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"你好\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"qwen-turbo\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"世界\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_qwen(&server); + let req = ChatRequest::builder("qwen-turbo", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 2); + let mut content = String::new(); + content.push_str(chunks[0].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "你好世界"); + } + + #[tokio::test] + async fn qwen_list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "qwen-turbo", "object": "model", "created": 1, "owned_by": "dashscope"}, + {"id": "qwen-vl-max", "object": "model", "created": 1, "owned_by": "dashscope"} + ] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_qwen(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "qwen-turbo"); + assert_eq!(models[0].provider, "qwen"); + } + + #[tokio::test] + async fn qwen_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "Invalid API key"}}).to_string()) + .create_async() + .await; + let adapter = make_qwen(&server); + let req = ChatRequest::builder("qwen-turbo", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn qwen_image_generate_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_qwen(&server); + let req = ImageRequest::builder("qwen-vl-max", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== Zhipu 元信息与能力 ==================== + + #[tokio::test] + async fn zhipu_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_zhipu(&server); + assert_eq!(adapter.provider_type(), "zhipu"); + assert_eq!(adapter.provider_name(), "智谱 AI"); + } + + #[tokio::test] + async fn zhipu_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_zhipu(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + } + + #[tokio::test] + async fn zhipu_chat_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_zhipu(&server); + let req = ChatRequest::builder("glm-4", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + } + + #[tokio::test] + async fn zhipu_list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [{"id": "glm-4", "object": "model", "created": 1, "owned_by": "zhipu"}] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_zhipu(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "glm-4"); + assert_eq!(models[0].provider, "zhipu"); + } + + #[tokio::test] + async fn zhipu_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_zhipu(&server); + let req = ChatRequest::builder("glm-4", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ==================== Doubao 元信息与能力 ==================== + + #[tokio::test] + async fn doubao_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_doubao(&server); + assert_eq!(adapter.provider_type(), "doubao"); + assert_eq!(adapter.provider_name(), "豆包"); + } + + #[tokio::test] + async fn doubao_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_doubao(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + } + + #[tokio::test] + async fn doubao_chat_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_doubao(&server); + let req = ChatRequest::builder("doubao-pro", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + } + + #[tokio::test] + async fn doubao_list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [{"id": "doubao-pro", "object": "model", "created": 1, "owned_by": "volcengine"}] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_doubao(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "doubao-pro"); + assert_eq!(models[0].provider, "doubao"); + } + + #[tokio::test] + async fn doubao_chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + let adapter = make_doubao(&server); + let req = ChatRequest::builder("doubao-pro", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + // ==================== Kimi 元信息与能力 ==================== + + #[tokio::test] + async fn kimi_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_kimi(&server); + assert_eq!(adapter.provider_type(), "kimi"); + assert_eq!(adapter.provider_name(), "Kimi (月之暗面)"); + } + + #[tokio::test] + async fn kimi_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_kimi(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + } + + #[tokio::test] + async fn kimi_chat_success() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_kimi(&server); + let req = ChatRequest::builder("moonshot-v1-8k", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + } + + #[tokio::test] + async fn kimi_list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [{"id": "moonshot-v1-8k", "object": "model", "created": 1, "owned_by": "moonshot"}] + }); + server + .mock("GET", "/models") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_kimi(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "moonshot-v1-8k"); + assert_eq!(models[0].provider, "kimi"); + } + + #[tokio::test] + async fn kimi_video_create_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_kimi(&server); + let req = VideoRequest::builder("moonshot-v1-8k", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== MiniMax 元信息与能力 ==================== + + #[tokio::test] + async fn minimax_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_minimax(&server); + assert_eq!(adapter.provider_type(), "minimax"); + assert_eq!(adapter.provider_name(), "MiniMax"); + } + + #[tokio::test] + async fn minimax_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_minimax(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn minimax_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_minimax(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + } + + #[tokio::test] + async fn minimax_chat_success() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let adapter = make_minimax(&server); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]) + .temperature(0.5) + .max_tokens(50) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + mock.assert_async().await; + } + + #[tokio::test] + async fn minimax_chat_stream_parses_sse() { + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"abab6.5-chat\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"hi\"},\"finish_reason\":null}]}\n\ + data: {\"id\":\"c1\",\"object\":\"chat.completion.chunk\",\"created\":1,\"model\":\"abab6.5-chat\",\"choices\":[{\"index\":0,\"delta\":{\"content\":\"!\"},\"finish_reason\":\"stop\"}]}\n\ + data: [DONE]\n"; + server + .mock("POST", "/chat/completions") + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_minimax(&server); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 2); + let mut content = String::new(); + content.push_str(chunks[0].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "hi!"); + } + + #[tokio::test] + async fn minimax_list_models_hardcoded() { + // MiniMax 无 /models 端点,list_models 返回硬编码列表,不应发请求 + let server = Server::new_async().await; + let adapter = make_minimax(&server); + let models = adapter.list_models(None).await.unwrap(); + assert!(!models.is_empty()); + // 验证含 chat 与 audio 两种类型 + assert!(models.iter().any(|m| m.model_type == ModelType::Chat)); + assert!(models.iter().any(|m| m.model_type == ModelType::Audio)); + // 验证 provider 字段 + assert!(models.iter().all(|m| m.provider == "minimax")); + } + + #[tokio::test] + async fn minimax_list_models_filter_by_type() { + let server = Server::new_async().await; + let adapter = make_minimax(&server); + let chat_models = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat_models.iter().all(|m| m.model_type == ModelType::Chat)); + assert!(chat_models.len() >= 5); + } + + #[tokio::test] + async fn minimax_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + let adapter = make_minimax(&server); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn minimax_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(429) + .with_body(json!({"error": {"message": "slow down"}}).to_string()) + .create_async() + .await; + let adapter = make_minimax(&server); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn minimax_chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/chat/completions") + .with_status(500) + .with_body(json!({"base_resp": {"status_msg": "internal error"}}).to_string()) + .create_async() + .await; + let adapter = make_minimax(&server); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + // base_resp.status_msg 应被提取为 message + assert!(message.contains("internal error")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn minimax_group_id_extracted_from_config() { + // group_id 经 config.extra["group_id"] 传入 + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .extra("group_id", "grp-123") + .build(); + let config = ProviderConfig::from_options("minimax", opts); + let adapter = MiniMaxAdapter::new(config).unwrap(); + assert_eq!(adapter.group_id(), Some("grp-123")); + } + + #[tokio::test] + async fn minimax_group_id_none_when_not_configured() { + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("minimax", opts); + let adapter = MiniMaxAdapter::new(config).unwrap(); + assert_eq!(adapter.group_id(), None); + } + + #[tokio::test] + async fn minimax_group_id_sent_as_header_when_configured() { + // 配置 group_id 时,chat 请求应携带 X-Group-Id header + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/chat/completions") + .match_header("x-group-id", "grp-123") + .with_status(200) + .with_body(openai_chat_body().to_string()) + .create_async() + .await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .extra("group_id", "grp-123") + .build(); + let config = ProviderConfig::from_options("minimax", opts); + let adapter = MiniMaxAdapter::new(config).unwrap(); + let req = ChatRequest::builder("abab6.5-chat", vec![ChatMessage::user("hi")]).build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn minimax_image_generate_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_minimax(&server); + let req = ImageRequest::builder("abab6.5-chat", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== Ernie 元信息与能力 ==================== + + #[tokio::test] + async fn ernie_provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_ernie(&server); + assert_eq!(adapter.provider_type(), "ernie"); + assert_eq!(adapter.provider_name(), "文心一言"); + } + + #[tokio::test] + async fn ernie_requires_api_key() { + let server = Server::new_async().await; + let adapter = make_ernie(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn ernie_capabilities_include_chat_and_vision() { + let server = Server::new_async().await; + let adapter = make_ernie(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + } + + #[tokio::test] + async fn ernie_chat_success_parses_result_field() { + // ERNIE 响应用 result 字段(非 OpenAI choices 结构) + let mut server = Server::new_async().await; + let body = json!({ + "id": "ernie-1", + "object": "chat.completion", + "created": 1700000000, + "result": "你好,我是文心一言", + "usage": {"prompt_tokens": 3, "completion_tokens": 8, "total_tokens": 11}, + "is_truncated": false + }); + let mock = server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/completions_pro".to_string(), + ), + ) + .match_query(mockito::Matcher::UrlEncoded( + "access_token".into(), + "test-token".into(), + )) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions_pro", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + assert_eq!(resp.id, "ernie-1"); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("你好,我是文心一言") + ); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(resp.usage.as_ref().unwrap().total_tokens, 11); + mock.assert_async().await; + } + + #[tokio::test] + async fn ernie_chat_truncated_returns_length_finish_reason() { + // is_truncated=true 时 finish_reason 应为 "length" + let mut server = Server::new_async().await; + let body = json!({ + "id": "ernie-2", + "result": "被截断的回复", + "is_truncated": true + }); + server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/completions".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions", vec![ChatMessage::user("hi")]).build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("length")); + } + + #[tokio::test] + async fn ernie_chat_extracts_system_to_separate_field() { + // system 消息应被提取到 system 字段(非 messages 数组) + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .match_body(mockito::Matcher::PartialJson(json!({ + "messages": [{"role": "user", "content": "hi"}], + "system": "你是助手" + }))) + .with_status(200) + .with_body(json!({"result": "ok"}).to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder( + "completions_pro", + vec![ChatMessage::system("你是助手"), ChatMessage::user("hi")], + ) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn ernie_chat_passes_temperature_and_top_p() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .match_body(mockito::Matcher::PartialJson(json!({ + "temperature": 0.8, + "top_p": 0.9 + }))) + .with_status(200) + .with_body(json!({"result": "ok"}).to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions_pro", vec![ChatMessage::user("hi")]) + .temperature(0.8) + .top_p(0.9) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn ernie_chat_stream_parses_sse_with_is_end() { + // ERNIE 流式:result 增量 + is_end 结束标记(无 [DONE]) + let mut server = Server::new_async().await; + let sse = "data: {\"id\":\"e1\",\"result\":\"你好\",\"is_end\":false}\n\ + data: {\"id\":\"e1\",\"result\":\"世界\",\"is_end\":false}\n\ + data: {\"id\":\"e1\",\"result\":\"\",\"is_end\":true,\"usage\":{\"prompt_tokens\":2,\"completion_tokens\":4,\"total_tokens\":6}}\n"; + server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(sse) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions_pro", vec![ChatMessage::user("hi")]).build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + assert_eq!(chunks.len(), 3); + // 前两块为增量文本,finish_reason 为 None + assert_eq!(chunks[0].choices[0].delta.content.as_deref(), Some("你好")); + assert!(chunks[0].choices[0].finish_reason.is_none()); + assert_eq!(chunks[1].choices[0].delta.content.as_deref(), Some("世界")); + // 第三块 is_end=true,finish_reason 为 "stop",并携带 usage + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(chunks[2].usage.as_ref().unwrap().total_tokens, 6); + } + + #[tokio::test] + async fn ernie_list_models_hardcoded() { + // ERNIE 无标准 /models 端点,list_models 返回硬编码列表 + let server = Server::new_async().await; + let adapter = make_ernie(&server); + let models = adapter.list_models(None).await.unwrap(); + assert!(!models.is_empty()); + assert!(models.iter().all(|m| m.provider == "ernie")); + // 验证含 ERNIE 4.0 / 3.5 + assert!(models.iter().any(|m| m.id == "completions_pro")); + assert!(models.iter().any(|m| m.id == "completions")); + } + + #[tokio::test] + async fn ernie_chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error_msg": "Invalid access_token"}).to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions_pro", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Authentication { message } => { + // error_msg 字段应被提取 + assert!(message.contains("Invalid access_token")); + } + _ => panic!("应为 Authentication"), + } + } + + #[tokio::test] + async fn ernie_chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock( + "POST", + mockito::Matcher::Regex( + r"^/rpc/2\.0/ai_custom/v1/wenxinworkshop/chat/".to_string(), + ), + ) + .match_query(mockito::Matcher::Any) + .with_status(429) + .with_body(json!({"error_msg": "qps limit"}).to_string()) + .create_async() + .await; + let adapter = make_ernie(&server); + let req = ChatRequest::builder("completions_pro", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn ernie_image_generate_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_ernie(&server); + let req = ImageRequest::builder("ernie-vilg", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn ernie_video_create_returns_unsupported() { + let server = Server::new_async().await; + let adapter = make_ernie(&server); + let req = VideoRequest::builder("ernie", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== Ernie api_key 解析与 oauth 流程 ==================== + + #[test] + fn ernie_new_with_plain_api_key_uses_as_access_token() { + // 无冒号的 api_key 直接作为 access_token,不触发 oauth + let opts = ClientOptions::builder() + .api_key("plain-token") + .base_url("https://aip.baidubce.com") + .build(); + let config = ProviderConfig::from_options("ernie", opts); + let adapter = ErnieAdapter::new(config).unwrap(); + assert!(!adapter.needs_oauth()); + assert_eq!(adapter.access_token(), Some("plain-token")); + } + + #[test] + fn ernie_new_with_ak_sk_format_marks_needs_oauth() { + // ak:sk 格式应拆分,标记需要 oauth,access_token 初始为 None + let opts = ClientOptions::builder() + .api_key("my-ak:my-sk") + .base_url("https://aip.baidubce.com") + .build(); + let config = ProviderConfig::from_options("ernie", opts); + let adapter = ErnieAdapter::new(config).unwrap(); + assert!(adapter.needs_oauth()); + assert_eq!(adapter.access_token(), None); + } + + #[tokio::test] + async fn ernie_start_fetches_access_token_when_ak_sk_configured() { + // 配置 ak:sk 时,start() 应调 /oauth/2.0/token 换取 access_token + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/oauth/2.0/token") + .match_query(mockito::Matcher::AllOf(vec![ + mockito::Matcher::UrlEncoded("grant_type".into(), "client_credentials".into()), + mockito::Matcher::UrlEncoded("client_id".into(), "my-ak".into()), + mockito::Matcher::UrlEncoded("client_secret".into(), "my-sk".into()), + ])) + .with_status(200) + .with_body(json!({"access_token": "fetched-token", "expires_in": 2592000}).to_string()) + .create_async() + .await; + + let opts = ClientOptions::builder() + .api_key("my-ak:my-sk") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("ernie", opts); + let mut adapter = ErnieAdapter::new(config).unwrap(); + assert!(adapter.needs_oauth()); + assert_eq!(adapter.access_token(), None); + + adapter.start().await.expect("start 应成功"); + assert_eq!(adapter.access_token(), Some("fetched-token")); + mock.assert_async().await; + } + + #[tokio::test] + async fn ernie_start_skips_oauth_when_plain_token() { + // 无冒号的 api_key,start() 不应发 oauth 请求 + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("plain-token") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("ernie", opts); + let mut adapter = ErnieAdapter::new(config).unwrap(); + adapter.start().await.unwrap(); + assert_eq!(adapter.access_token(), Some("plain-token")); + } + + #[tokio::test] + async fn ernie_oauth_error_propagates() { + // oauth 端点返回错误时,start() 应传播错误 + let mut server = Server::new_async().await; + server + .mock("GET", "/oauth/2.0/token") + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body(json!({"error": "invalid client"}).to_string()) + .create_async() + .await; + let opts = ClientOptions::builder() + .api_key("bad-ak:bad-sk") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("ernie", opts); + let mut adapter = ErnieAdapter::new(config).unwrap(); + let err = adapter.start().await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ==================== 错误映射单元测试 ==================== + + #[test] + fn map_ernie_error_extracts_error_msg_field() { + let body = json!({"error_msg": "some error"}).to_string(); + let err = map_ernie_error(500, &body); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "some error"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_ernie_error_falls_back_to_message_field() { + let body = json!({"message": "fallback msg"}).to_string(); + let err = map_ernie_error(400, &body); + match err { + AibridgeError::Validation { message, .. } => assert_eq!(message, "fallback msg"), + _ => panic!("应为 Validation"), + } + } + + #[test] + fn map_ernie_error_no_json_falls_back_to_http_status() { + let err = map_ernie_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_minimax_error_extracts_base_resp_status_msg() { + // MiniMax 特有:base_resp.status_msg 字段 + let body = json!({"base_resp": {"status_msg": "minimax error"}}).to_string(); + let err = map_minimax_error(500, &body); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "minimax error"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_minimax_error_extracts_error_message_first() { + // 优先 error.message,而非 base_resp.status_msg + let body = json!({ + "error": {"message": "primary"}, + "base_resp": {"status_msg": "secondary"} + }) + .to_string(); + let err = map_minimax_error(500, &body); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "primary"), + _ => panic!("应为 Api"), + } + } + + // ==================== Ernie 消息转换单元测试 ==================== + + #[test] + fn ernie_convert_messages_extracts_system() { + let messages = vec![ + ChatMessage::system("sys-prompt"), + ChatMessage::user("hello"), + ChatMessage::assistant("hi"), + ]; + let (msgs, system) = ErnieAdapter::convert_messages(&messages); + assert_eq!(system.as_deref(), Some("sys-prompt")); + assert_eq!(msgs.len(), 2); + assert_eq!(msgs[0]["role"], "user"); + assert_eq!(msgs[0]["content"], "hello"); + assert_eq!(msgs[1]["role"], "assistant"); + assert_eq!(msgs[1]["content"], "hi"); + } + + #[test] + fn ernie_convert_messages_no_system() { + let messages = vec![ChatMessage::user("hello")]; + let (msgs, system) = ErnieAdapter::convert_messages(&messages); + assert!(system.is_none()); + assert_eq!(msgs.len(), 1); + assert_eq!(msgs[0]["role"], "user"); + } + + // ==================== start / close 是空操作(OpenAI 兼容族) ==================== + + #[tokio::test] + async fn qwen_start_and_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_qwen(&server); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[tokio::test] + async fn minimax_start_and_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_minimax(&server); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } +} From 38801686c93612a9b87ced8616dcef8533d4a57f Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 21:44:24 +0800 Subject: [PATCH 31/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2a=20=E7=AC=AC=E4=B8=89=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20emerging=5Fmodels=20+=20chinese=20=E5=88=B0?= =?UTF-8?q?=E5=B7=A5=E5=8E=82=EF=BC=88=E9=98=B6=E6=AE=B52a=20=E5=AE=8C?= =?UTF-8?q?=E6=88=90=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 123 ++++++++++++++++++-- crates/aibridge-core/src/adapters/mod.rs | 6 + 2 files changed, 122 insertions(+), 7 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index a87efec..eaf9a00 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -16,7 +16,11 @@ use crate::adapters::aggregation_platforms::{ }; use crate::adapters::agnes::AgnesAdapter; use crate::adapters::azure::AzureAdapter; +use crate::adapters::chinese::{ + DoubaoAdapter, ErnieAdapter, KimiAdapter, MiniMaxAdapter, QwenAdapter, ZhipuAdapter, +}; use crate::adapters::echo::EchoAdapter; +use crate::adapters::emerging_models::{IdeogramAdapter, LlamaAdapter, LumaAdapter}; use crate::adapters::gemini::GeminiAdapter; use crate::adapters::more_models::{ CohereAdapter, DeepSeekAdapter, MistralAdapter, PerplexityAdapter, StepFunAdapter, @@ -52,13 +56,23 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "mistral", "cohere", "perplexity", + // 阶段 2a 第三批 emerging_models: + "ideogram", + "luma", + "llama", + // 阶段 2a 第三批 chinese: + "qwen", + "zhipu", + "doubao", + "ernie", + "kimi", + "minimax", // 阶段 2b/2c 待实现: "anthropic", "runway", "pika", "kling", "stability", - "chinese", "edge-tts", "elevenlabs", "cartesia", @@ -102,13 +116,23 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { "mistral" => Ok(Box::new(MistralAdapter::new(config)?)), "cohere" => Ok(Box::new(CohereAdapter::new(config)?)), "perplexity" => Ok(Box::new(PerplexityAdapter::new(config)?)), + // 新兴模型:别名对齐 Python agn/adapters/emerging_models.py 末尾 register 调用 + "ideogram" | "ideo" => Ok(Box::new(IdeogramAdapter::new(config)?)), + "luma" | "dream-machine" | "lumalabs" => Ok(Box::new(LumaAdapter::new(config)?)), + "llama" | "meta-llama" | "meta" => Ok(Box::new(LlamaAdapter::new(config)?)), + // 中文模型:别名对齐 Python agn/adapters/chinese.py 末尾 register 调用 + // (qwen/zhipu/doubao/ernie/kimi/minimax 均无别名,直接注册主名) + "qwen" => Ok(Box::new(QwenAdapter::new(config)?)), + "zhipu" => Ok(Box::new(ZhipuAdapter::new(config)?)), + "doubao" => Ok(Box::new(DoubaoAdapter::new(config)?)), + "ernie" => Ok(Box::new(ErnieAdapter::new(config)?)), + "kimi" => Ok(Box::new(KimiAdapter::new(config)?)), + "minimax" => Ok(Box::new(MiniMaxAdapter::new(config)?)), // 阶段 2 适配器占位 - "anthropic" | "runway" | "pika" | "kling" | "stability" | "chinese" | "edge-tts" - | "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { - Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2 待实现)"), - }) - } + "anthropic" | "runway" | "pika" | "kling" | "stability" | "edge-tts" | "elevenlabs" + | "cartesia" | "deepgram" | "assemblyai" => Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 2 待实现)"), + }), // 未知 provider _ => Err(AibridgeError::provider_not_found(format!( "{provider}(未知 provider,支持:{})", @@ -301,6 +325,81 @@ mod tests { assert_eq!(adapter.provider_type(), "perplexity"); } + #[test] + fn create_ideogram_returns_adapter() { + // 阶段 2a emerging_models:IdeogramAdapter 自带 DEFAULT_IDEOGRAM_BASE_URL 兜底,仅需 api_key + let adapter = create_adapter(config_for("ideogram")).expect("工厂应能创建 ideogram 适配器"); + assert_eq!(adapter.provider_type(), "ideogram"); + } + + #[test] + fn create_luma_returns_adapter() { + let adapter = create_adapter(config_for("luma")).expect("工厂应能创建 luma 适配器"); + assert_eq!(adapter.provider_type(), "luma"); + } + + #[test] + fn create_llama_returns_adapter() { + let adapter = create_adapter(config_for("llama")).expect("工厂应能创建 llama 适配器"); + assert_eq!(adapter.provider_type(), "llama"); + } + + #[test] + fn create_emerging_models_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/emerging_models.py 末尾 register 调用: + // ideo -> ideogram / dream-machine & lumalabs -> luma / meta-llama & meta -> llama + let ideo = create_adapter(config_for("ideo")).expect("别名 ideo 应映射到 ideogram"); + assert_eq!(ideo.provider_type(), "ideogram"); + let dream_machine = + create_adapter(config_for("dream-machine")).expect("别名 dream-machine 应映射到 luma"); + assert_eq!(dream_machine.provider_type(), "luma"); + let lumalabs = create_adapter(config_for("lumalabs")).expect("别名 lumalabs 应映射到 luma"); + assert_eq!(lumalabs.provider_type(), "luma"); + let meta_llama = + create_adapter(config_for("meta-llama")).expect("别名 meta-llama 应映射到 llama"); + assert_eq!(meta_llama.provider_type(), "llama"); + let meta = create_adapter(config_for("meta")).expect("别名 meta 应映射到 llama"); + assert_eq!(meta.provider_type(), "llama"); + } + + #[test] + fn create_qwen_returns_adapter() { + // 阶段 2a chinese:QwenAdapter 自带 DEFAULT_QWEN_BASE_URL 兜底,仅需 api_key + let adapter = create_adapter(config_for("qwen")).expect("工厂应能创建 qwen 适配器"); + assert_eq!(adapter.provider_type(), "qwen"); + } + + #[test] + fn create_zhipu_returns_adapter() { + let adapter = create_adapter(config_for("zhipu")).expect("工厂应能创建 zhipu 适配器"); + assert_eq!(adapter.provider_type(), "zhipu"); + } + + #[test] + fn create_doubao_returns_adapter() { + let adapter = create_adapter(config_for("doubao")).expect("工厂应能创建 doubao 适配器"); + assert_eq!(adapter.provider_type(), "doubao"); + } + + #[test] + fn create_ernie_returns_adapter() { + // ErnieAdapter: api_key 不含 ':' 时直接作 access_token,构造无特殊要求 + let adapter = create_adapter(config_for("ernie")).expect("工厂应能创建 ernie 适配器"); + assert_eq!(adapter.provider_type(), "ernie"); + } + + #[test] + fn create_kimi_returns_adapter() { + let adapter = create_adapter(config_for("kimi")).expect("工厂应能创建 kimi 适配器"); + assert_eq!(adapter.provider_type(), "kimi"); + } + + #[test] + fn create_minimax_returns_adapter() { + let adapter = create_adapter(config_for("minimax")).expect("工厂应能创建 minimax 适配器"); + assert_eq!(adapter.provider_type(), "minimax"); + } + #[test] fn create_additional_models_aliases_map_to_main_provider_type() { // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: @@ -357,6 +456,16 @@ mod tests { assert!(is_known_provider("mistral")); assert!(is_known_provider("cohere")); assert!(is_known_provider("perplexity")); + // 阶段 2a 第三批 emerging_models + chinese + assert!(is_known_provider("ideogram")); + assert!(is_known_provider("luma")); + assert!(is_known_provider("llama")); + assert!(is_known_provider("qwen")); + assert!(is_known_provider("zhipu")); + assert!(is_known_provider("doubao")); + assert!(is_known_provider("ernie")); + assert!(is_known_provider("kimi")); + assert!(is_known_provider("minimax")); } #[test] diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 0a7a1e5..c3f25f1 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -39,3 +39,9 @@ pub mod additional_models; /// 更多模型适配器:阶段 2a,含 DeepSeek/StepFun/Mistral/Cohere/Perplexity 五个 OpenAI 兼容子适配器 pub mod more_models; + +/// 新兴模型适配器:阶段 2a,含 Ideogram/Luma/Llama 三个子适配器 +pub mod emerging_models; + +/// 中文模型适配器:阶段 2a,含 Qwen/Zhipu/Doubao/Ernie/Kimi/MiniMax 六个 OpenAI 兼容子适配器 +pub mod chinese; From 37a1b9a7adb7d111c0bdefb597add40e2541a3d3 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Tue, 7 Jul 2026 22:28:51 +0800 Subject: [PATCH 32/55] =?UTF-8?q?docs:=20AIBridge=20Rust=20=E9=87=8D?= =?UTF-8?q?=E6=9E=84=E8=BF=9B=E5=BA=A6=E6=96=87=E6=A1=A3=20+=20AGENTS.md?= =?UTF-8?q?=20=E6=8E=A5=E6=89=8B=E6=8C=87=E5=BC=95?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 新增 docs/PROGRESS.md:完整进度(阶段0-2a 完成,23 provider,810 单测)+ 接手指南(环境、命令、worktree 模式、阶段 2b/2c 要点、关键约束)+ 提交历史,供其他 agent 接手。 AGENTS.md 顶部加 Rust 重构指向提示。 --- AGENTS.md | 4 + docs/PROGRESS.md | 217 +++++++++++++++++++++++++++++++++++++++++++++++ 2 files changed, 221 insertions(+) create mode 100644 docs/PROGRESS.md diff --git a/AGENTS.md b/AGENTS.md index 8474a4c..01928ac 100644 --- a/AGENTS.md +++ b/AGENTS.md @@ -2,6 +2,10 @@ > 本文件是 Agent 执行任务时必须首先加载和理解的项目指南,包含项目目标、架构、规范、验收标准等关键信息。 +> ⚠️ **Rust 重构进行中**:本项目正在用 Rust 重构为跨语言 SDK `aibridge`(分支 `feat/aibridge-rust-rewrite`)。 +> 接手 Rust 重构工作请先读 **[docs/PROGRESS.md](docs/PROGRESS.md)**(完整进度 + 接手指南)+ [设计文档](docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md) + [实现计划](docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md)。 +> 本文档(AGENTS.md)描述的是 Python 旧版(agn-sdk v1),Rust 新版以 PROGRESS.md + 设计文档为准。 + --- ## 1. 项目概述 diff --git a/docs/PROGRESS.md b/docs/PROGRESS.md new file mode 100644 index 0000000..7b1b30d --- /dev/null +++ b/docs/PROGRESS.md @@ -0,0 +1,217 @@ +# AIBridge Rust 重构 · 进度与接手文档 + +> 本文档供任何 agent 接手 AIBridge Rust 重构工作使用。自包含,不依赖 Claude memory。 +> 最后更新:2026-07-07 + +--- + +## 1. 项目概述 + +将 Python `agn-sdk`(多模态 AI 统一接口 SDK,v1.3.3,~19700 行)用 Rust 重构为跨语言 SDK `aibridge`,支持五种语言直接 import。 + +- **分支**:`feat/aibridge-rust-rewrite`(基于 `main`,`main` 仍是 Python 旧版) +- **品牌名**:aibridge(原 agn-sdk,PyPI/npm 均可用) +- **五语言**:Python(PyO3 直连 core)/ JS-TS(napi-rs 直连 core)/ Go(CGO 调 ffi)/ JVM(JNA 调 ffi)/ .NET(P/Invoke 调 ffi) + +## 2. 关键文档(必读) + +| 文档 | 路径 | 内容 | +|---|---|---| +| 设计文档 | [docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md](superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md) | 架构、数据模型、FFI 边界、异步桥接、错误处理、适配器迁移策略 | +| 实现计划 | [docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md](superpowers/plans/2026-07-07-aibridge-implementation-plan.md) | 阶段 0-3 任务分解、多 agent 编排策略、里程碑 | +| 本进度文档 | docs/PROGRESS.md | 当前进度 + 接手指南(本文档) | + +## 3. 架构概览 + +``` + aibridge-core (Rust, 纯 async 逻辑) + ┌──────────┴──────────┐ + 直连(原生async) C ABI (aibridge-ffi cdylib) + ┌─────┴─────┐ ┌─────┬─────┬─────┐ + aibridge- aibridge- aibridge- aibridge- aibridge- + python node go jvm dotnet + (PyO3) (napi-rs) (CGO) (JNA) (P/Invoke) +``` + +### Monorepo 布局 +``` +aibridge/ +├── crates/ +│ ├── aibridge-core/ # Rust 核心(纯逻辑,无 FFI) +│ ├── aibridge-ffi/ # C ABI cdylib(给 Go/JVM/.NET) +│ ├── aibridge-python/ # PyO3 绑定(直连 core) +│ └── aibridge-node/ # napi-rs 绑定(直连 core) +├── bindings/ +│ ├── go/ # CGO 调 ffi +│ ├── jvm/ # JNA 调 ffi(Java) +│ └── dotnet/ # P/Invoke 调 ffi(C#) +├── docs/ # 设计文档 + 计划 + 本进度文档 +└── examples/ # 五语言 hello world(echo adapter) +``` + +## 4. 当前进度(截至 2026-07-07) + +### ✅ 阶段 0:地基(完成) +- Cargo workspace + 4 crate 骨架 +- aibridge-core:error/config/http/retry/model/adapter/client/router +- aibridge-ffi:C ABI(全局 tokio runtime + 句柄 + JSON 边界 + cbindgen 生成 include/aibridge.h) +- 五语言绑定 hello world 跑通(Python/Node/Go/JVM,.NET 代码就绪待 dotnet sdk) + +### ✅ 阶段 1:MVP 四 provider(完成) +- openai / agnes / 火山(volcengine_cv) / gemini +- 五语言可用(含流式 chat_stream) +- Python/Node 流式已重构(PyO3 Coroutine / napi async fn,真实 IO 不阻塞) + +### ✅ 阶段 1.5:质量收尾(完成) +- 跨语言一致性测试(四语言 chat/stream/speech/错误全一致) +- 错误 code 统一(对齐 Rust error.rs 的 code() 值) +- dylib 产物名冲突修复(aibridge-python lib 改 `_aibridge`) + +### ✅ 阶段 2a:OpenAI 兼容族 + 部分独立协议(完成,23 provider) +| 文件 | provider(含别名) | +|---|---| +| openai.rs | openai | +| agnes.rs | agnes | +| volcengine_cv.rs | volcengine_cv(火山引擎) | +| gemini.rs | gemini | +| azure.rs | azure | +| aggregation_platforms.rs | siliconflow\|sf, togetherai\|together, fireworksai\|fireworks, cloudflareai\|cloudflare\|workersai | +| additional_models.rs | grok\|xaigrok, yi\|lingyiwanwu, sensenova\|shangtang, hunyuan\|tencent_hunyuan, groq | +| more_models.rs | deepseek, stepfun\|step, mistral, cohere, perplexity | +| emerging_models.rs | ideogram\|ideo, luma\|dream-machine\|lumalabs, llama\|meta-llama\|meta | +| chinese.rs | qwen, zhipu, doubao, ernie, kimi, minimax | + +### ⏳ 阶段 2b:独立协议 5 个(待做) +anthropic / stability / runway / pika / kling + +### ⏳ 阶段 2c:音频 5 个(待做,二进制载荷) +edge-tts / elevenlabs / cartesia / deepgram / assemblyai + +### ⏳ 阶段 3:发布(待做) +- CI 矩阵(平台 × 语言,交叉编译) +- 五语言包发布(PyPI `aibridge` / npm `aibridge` / Maven `io.aibridge:aibridge` / NuGet `AIBridge` / Go module `aibridge-go`) +- Python v1→v2 迁移指南 +- 文档网站 +- 旧版 v1 归档 + 打 v2.0.0 tag + +## 5. 测试状态 + +- **aibridge-core**:810 单测全通过 +- **aibridge-ffi**:39 单测全通过 +- **五语言 hello world**(echo adapter):Python/Node/Go/JVM 跑通,.NET 代码就绪待 dotnet +- **跨语言一致性测试**:tests/consistency/(四语言 chat/stream/speech/错误全一致) + +## 6. 遗留问题(非阻塞) + +1. **.NET hello world**:待装 dotnet sdk(`brew install --cask dotnet-sdk`,需 sudo),代码 + P/Invoke 已就绪 +2. **一致性测试纳入 CI**:当前手动跑,待接入 CI matrix +3. **真实 IO 流式验证**:Python/Node 流式重构已完成(不阻塞推理),但真实 API key 验证待用户 +4. **dylib 分发**:Go/JVM/.NET 依赖 libaibridge,JVM/.NET 打进包,Go 提供安装脚本(阶段 3 处理) + +## 7. 新 agent 接手指南 + +### 7.1 环境 +- Rust 1.94 + cargo(crates.io 镜像已配 `~/.cargo/config.toml` 走 rsproxy,国内网络必须) +- Python 3.14 + maturin(`pip install maturin`) +- Node 25 + npm +- Go 1.26 +- Java 21 (Temurin) +- dotnet 未装(.NET 绑定代码就绪,hello world 待验证) + +### 7.2 验证当前状态 +```bash +git checkout feat/aibridge-rust-rewrite +cargo test -p aibridge-core # 期望 810 passed +cargo test -p aibridge-ffi # 期望 39 passed +cargo build --workspace # 0 warning +``` + +### 7.3 怎么继续阶段 2b/2c + +**模式**:每批 2 个适配器,worktree 并行实现 + 收尾 agent cherry-pick 注册(同阶段 2a)。 + +1. **启动 2 个 worktree agent**(每个实现一个 adapter .rs): + - prompt 要点:读设计文档第 10 节 + Python 对应 `agn/adapters/.py` + `openai_compat.rs`(地基)+ `volcengine_cv.rs`/`more_models.rs Cohere`(独立协议范例);实现 adapter .rs;mockito 单测;**只 git add 自己的 .rs**(不 add mod.rs/factory.rs);commit 返回 hash + - `isolation: "worktree"` + `run_in_background: true` + - worktree 初始可能在 main 分支(无 crates),需 `git checkout feat/aibridge-rust-rewrite` 或基于它建工作分支 +2. **收尾 agent**:cherry-pick 2 个 commit + 注册 mod.rs(pub mod)+ factory.rs(match 分支 + 别名,参考 Python `agn/adapters/factory.py` 的 register)+ 更新 factory 测试 + `cargo test` 全量 + commit +3. **避免 4+ 并行**(会触发 API 限流 429),每批 2 个 + +### 7.4 阶段 2b 各适配器要点 +- **anthropic**:Claude messages API,`POST /messages`,header `x-api-key`+`anthropic-version`,流式 SSE 事件(content_block_delta),`ANTHROPIC_MAPPING`。Python: `agn/adapters/anthropic.py` +- **stability**:Stability AI 图像协议,独立。Python: `agn/adapters/stability.py` +- **runway**:视频协议,独立。Python: `agn/adapters/runway.py` +- **pika**:视频协议,独立。Python: `agn/adapters/pika.py` +- **kling**:可灵视频/图像,独立。Python: `agn/adapters/kling.py` + +### 7.5 阶段 2c 音频要点 +- 二进制载荷(TTS 返 audio_data bytes,ASR 接受 file path/URL/bytes/base64) +- edge-tts 免认证(`requires_api_key=false`,加到 `client.rs` 的 `is_free_provider`) +- TTS 音色健康检查/推荐/自动降级(v1.3.3 特性,保留) +- Python: `agn/adapters/audio_adapters.py`(含全部 5 个) + +### 7.6 关键约束(必须遵守) +- **中文注释**(项目强制规则,文件模块文档字符串 + 公开项文档注释) +- **错误 code 对齐** Rust `aibridge-core/src/error.rs` 的 `code()` 实际值(带 `_error` 后缀) +- **mockito 1.x 用法**:`let mut server = mockito::Server::new_async().await;`(不是 `async_try_start`) +- **aibridge-python lib name = `_aibridge`**(避免与 ffi 的 libaibridge dylib 冲突,不要改回 `aibridge`) +- **Python/Node 流式**已重构(PyO3 Coroutine / napi async fn),不要再改回 block_on +- **Adapter trait** 在 `crates/aibridge-core/src/adapter/base.rs`;工厂注册在 `adapter/factory.rs`(编译期 match,非运行时注册) +- **openai_compat.rs** 的 9 个方法已 pub,子适配器组合委托复用 + +### 7.7 关键命令 +```bash +# Rust 核心 +cargo test -p aibridge-core +cargo build --workspace +cargo clippy -p aibridge-core -- -D warnings + +# Python 绑定 +pip install maturin +maturin develop -m crates/aibridge-python/Cargo.toml +python examples/hello_python.py + +# Node 绑定 +cd crates/aibridge-node && npm install && napi build && cd ../.. +node examples/hello_node.js + +# Go 绑定 +cargo build -p aibridge-ffi +cd bindings/go && CGO_ENABLED=1 DYLD_LIBRARY_PATH=../../target/debug go run ./example + +# JVM 绑定 +cd bindings/jvm && ./gradlew run +``` + +## 8. 提交历史(阶段 0-2a) + +``` +3880168 feat(aibridge-core): 阶段2a 第三批收尾(emerging_models + chinese,阶段2a 完成) +a0eb96d feat(aibridge-core): 阶段2a 第二批收尾(additional_models + more_models) +cb4536a feat(aibridge-core): 阶段2a 第一批收尾(azure + 聚合平台) +7b4a79d fix: 阶段1 Python/Node 流式桥接重构 +e5df4f3 fix(aibridge-python): dylib 产物名改为 _aibridge +68954d2 feat: 阶段1.5 跨语言一致性测试 + 错误 code 统一 +bf5cb1d feat(aibridge-core): 阶段1 收尾 注册四 MVP 适配器到工厂 +60de645 feat(aibridge-core): 阶段1.0 OpenAI 兼容适配器地基 +b81dab5 chore: 提交阶段0.6 五语言绑定依赖锁 +0a802b9 feat(aibridge-dotnet): 阶段0.6 P/Invoke 绑定 +268cd49 feat(aibridge-python): 阶段0.6 PyO3 绑定 + hello world +2ecfb58 feat(aibridge-node): 阶段0.6 napi-rs 绑定 + hello world +5118eeb feat(aibridge-jvm): 阶段0.6 JNA 绑定 + hello world +f36ee7c feat(aibridge-go): 阶段0.6 CGO 绑定 + hello world +8020c9c feat(aibridge-core): echo 适配器用于阶段0.6 管线验证 +ddafa26 feat(aibridge-ffi): 阶段0.5 C ABI 层 +414678d feat(aibridge-core): 阶段0.2-0.4 基础设施/数据模型/Adapter/Client/Router +1b7414e chore: 提交 Cargo.lock +350177c feat(aibridge): 阶段0.1 Cargo workspace 骨架 +65fcf2b docs: AIBridge Rust 重构设计文档 +``` + +## 9. 用户偏好(接手 agent 必读) + +- 用户是**编程小白**,沟通用**大白话**,少术语多类比 +- 以**行业专家身份替其做技术决策**,纯技术取舍直接拍板并解释理由 +- 只在**影响成本/时间/兼容性的业务分叉**处征求其意见 +- 用户**全权委托** Claude 端到端实施,实施阶段用多 agent 并行 +- 所有回复、代码注释、文档用**中文**(技术标识符除外) From 954d153757340947e6bea9482e135c4d9d4e203a Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 00:42:25 +0800 Subject: [PATCH 33/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20anthropic=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../aibridge-core/src/adapters/anthropic.rs | 1937 +++++++++++++++++ 1 file changed, 1937 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/anthropic.rs diff --git a/crates/aibridge-core/src/adapters/anthropic.rs b/crates/aibridge-core/src/adapters/anthropic.rs new file mode 100644 index 0000000..a374378 --- /dev/null +++ b/crates/aibridge-core/src/adapters/anthropic.rs @@ -0,0 +1,1937 @@ +//! Anthropic Claude 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/anthropic.py`。 +//! +//! Anthropic Claude 采用独立协议(非 OpenAI 兼容),本模块独立实现 chat / chat_stream / +//! list_models,不复用 `OpenAiCompatAdapter` 的请求体构造,但组合它以复用 HttpClient、 +//! 能力集合与错误映射(与 Cohere 适配器的组合模式一致)。 +//! +//! ## 协议要点 +//! +//! - Base URL: `https://api.anthropic.com/v1` +//! - Chat: `POST /messages`(流式同端点,body 中 `stream: true`) +//! - 认证: `x-api-key: {key}` header(非 Bearer) +//! - 版本: `anthropic-version: 2023-06-01` header(固定值) +//! - 请求体: `{model, messages, system, max_tokens, ...}` +//! - `system` 为顶层独立字段(从统一消息的 system 消息提取,不在 messages 中) +//! - `max_tokens` 必填(默认 1024) +//! - `stop` → `stop_sequences`(数组) +//! - `reasoning_effort` → `thinking`(对齐 Python ANTHROPIC_MAPPING 的 value_map) +//! - 响应: `content[].text` 拼接为回复文本;`stop_reason` 映射为 finish_reason; +//! `usage.input_tokens` / `usage.output_tokens` 映射为 prompt / completion tokens +//! - 流式: SSE,每行 `data: `,JSON 含 `type` 字段: +//! - `content_block_delta`(delta.type==text_delta)→ 增量文本 +//! - `message_delta`(delta.stop_reason)→ 结束原因 +//! - `message_stop` → 流结束 +//! - Models: `GET /models`,响应 `{"data":[{"id":...,"display_name":...}]}` +//! +//! 官方文档: + +use std::collections::HashMap; + +use async_trait::async_trait; +use futures::stream::{StreamExt, TryStreamExt}; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet, ChatStream}; +use crate::adapters::openai_compat::OpenAiCompatAdapter; +use crate::config::ProviderConfig; +use crate::error::{AibridgeError, Result}; +use crate::model::chat::{ + ChatChoice, ChatCompletion, ChatCompletionChunk, ChatCompletionDelta, ChatMessage, ChatRequest, + ChoiceMessage, ContentPart, DeltaMessage, UserContent, +}; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::options::{ParameterMapping, ReasoningEffort, StopSeq}; +use crate::util; + +// ==================== 默认配置 ==================== + +/// Anthropic 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL = "https://api.anthropic.com"`, +/// Rust 直接采用含 `/v1` 前缀的值(Python 在 httpx base_url 上拼接 `/v1/messages`, +/// 此处 base_url 已含 `/v1`,chat 端点为 `messages`,行为等价)。 +pub const DEFAULT_ANTHROPIC_BASE_URL: &str = "https://api.anthropic.com/v1"; + +/// Anthropic API 版本(固定值,对应 Python v1 `DEFAULT_API_VERSION`) +pub const ANTHROPIC_API_VERSION: &str = "2023-06-01"; + +/// 默认 max_tokens(Anthropic 必填字段,未指定时兜底,对齐 Python v1 默认 1024) +const DEFAULT_MAX_TOKENS: u32 = 1024; + +/// Anthropic 参数映射 +/// +/// 对应 Python v1 `ANTHROPIC_MAPPING`。 +/// Rust 的 `ParameterMapping` 仅支持 rename_map(无 value_map),故 `reasoning_effort → thinking` +/// 的值映射在 `build_chat_body` 中特殊处理(参照 DeepSeek 自动注入 thinking 的模式)。 +/// +/// rename_map: +/// - `stop` → `stop_sequences`(其余参数 max_tokens / top_p / top_k / temperature 原名透传) +pub fn anthropic_mapping() -> ParameterMapping { + let mut rename: HashMap> = HashMap::new(); + rename.insert("stop".into(), Some("stop_sequences".into())); + ParameterMapping { rename_map: rename } +} + +/// Anthropic 支持的能力集合 +/// +/// 对齐 Python v1 `supported_capabilities`(CHAT / CHAT_STREAM / VISION)。 +/// Rust 未声明 ToolCall:工具调用的完整请求体/响应解析未在本阶段实现 +/// (tool 消息按 user/tool_result 简化转换,保证不丢信息但不声明能力)。 +fn anthropic_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::Chat); + caps.insert(Capabilities::ChatStream); + caps.insert(Capabilities::Vision); + caps +} + +// ==================== 适配器结构 ==================== + +/// Anthropic Claude 适配器 +/// +/// 组合 `OpenAiCompatAdapter` 仅复用其 HttpClient、能力集合与错误映射, +/// chat / chat_stream / list_models 独立实现(Anthropic 特有 `/messages` 端点 + +/// `x-api-key` 认证 + `system` 顶层字段 + `content[].text` 响应 + SSE 事件)。 +/// +/// - Base URL: `https://api.anthropic.com/v1` +/// - Chat: `POST /messages` +/// - Models: `GET /models` +/// - 认证: `x-api-key` header + `anthropic-version` header +/// - 文档: +pub struct AnthropicAdapter { + /// OpenAI 兼容地基(仅复用 HttpClient + 错误映射 + 能力集合,不委托 chat 等方法) + compat: OpenAiCompatAdapter, +} + +impl AnthropicAdapter { + /// 创建 Anthropic 适配器 + pub fn new(config: ProviderConfig) -> Result { + let compat = OpenAiCompatAdapter::new( + config, + Self::PROVIDER_TYPE, + Self::PROVIDER_NAME, + DEFAULT_ANTHROPIC_BASE_URL, + anthropic_capabilities(), + )?; + Ok(Self { compat }) + } + + /// 用显式 compat 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_compat(compat: OpenAiCompatAdapter) -> Self { + Self { compat } + } + + /// Provider 类型标识 + const PROVIDER_TYPE: &'static str = "anthropic"; + + /// Provider 显示名称 + const PROVIDER_NAME: &'static str = "Anthropic Claude"; + + /// API key(从兼容地基配置提取) + fn api_key(&self) -> Option<&str> { + self.compat.api_key() + } + + /// base_url + fn base_url(&self) -> &str { + self.compat.base_url() + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url().trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验能力是否被支持 + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.compat.capabilities_set().contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: {})", cap.as_str(), Self::PROVIDER_TYPE), + }) + } + } + + /// 发送带认证的 POST JSON 请求,并用错误映射处理响应 + /// + /// Anthropic 用 `x-api-key` + `anthropic-version` header 认证(非 Bearer)。 + /// 错误映射复用 `OpenAiCompatAdapter::map_api_error`:其 `parse_error_message` + /// 会提取 Anthropic 错误体 `{"type":"error","error":{"message":"..."}}` 中的 + /// `error.message`,分类规则(401→Auth、429→RateLimit、404→ModelNotFound、 + /// 400→Validation、5xx→Api)与任务规格一致。 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self.send_post(&url, body).await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造并发送带 Anthropic 认证 header 的 POST 请求 + async fn send_post( + &self, + url: &str, + body: &Value, + ) -> std::result::Result { + self.compat + .http_inner() + .post(url) + .header("x-api-key", self.api_key().unwrap_or("")) + .header("anthropic-version", ANTHROPIC_API_VERSION) + .header("content-type", "application/json") + .json(body) + .send() + .await + } + + /// 发送带认证的 GET 请求,并用错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .compat + .http_inner() + .get(&url) + .header("x-api-key", self.api_key().unwrap_or("")) + .header("anthropic-version", ANTHROPIC_API_VERSION) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 将 API 错误响应映射为 AibridgeError + /// + /// 复用 `OpenAiCompatAdapter::map_api_error`:其错误消息提取兼容 Anthropic 的 + /// `error.message` 结构,状态码分类符合任务规格。 + fn map_api_error(status: u16, body: &str) -> AibridgeError { + OpenAiCompatAdapter::map_api_error(status, body) + } + + // ==================== 消息转换 ==================== + + /// 转换统一消息为 Anthropic 格式,并提取 system prompt + /// + /// 对齐 Python v1 `AnthropicAdapter._convert_messages`: + /// - system 消息 → 提取为顶层 `system` 字段(多条用 `\n` 拼接,Python 取最后一条; + /// Rust 改为拼接以保留全部 system 上下文) + /// - user 消息 → `{"role":"user","content":...}`(纯文本用字符串,多模态用 content blocks) + /// - assistant 消息 → `{"role":"assistant","content":...}`(tool_calls 暂未完整序列化为 + /// tool_use 块,本阶段聚焦文本对话) + /// - tool 消息 → `{"role":"user","content":[{"type":"tool_result",...}]}`(Anthropic 原生格式) + /// + /// 多模态图片: + /// - `data:` URI → `{"type":"image","source":{"type":"base64","media_type":...,"data":...}}` + /// - `http(s):` URL → `{"type":"image","source":{"type":"url","url":...}}`(Anthropic 新版支持) + fn convert_messages(messages: &[ChatMessage]) -> (Vec, Option) { + let mut converted: Vec = Vec::new(); + let mut system_parts: Vec = Vec::new(); + + for msg in messages { + match msg { + ChatMessage::System { content, .. } => { + system_parts.push(content.clone()); + } + ChatMessage::User { content, .. } => { + let c = match content { + UserContent::Text(s) => Value::String(s.clone()), + UserContent::Parts(parts) => { + let blocks: Vec = + parts.iter().filter_map(convert_content_part).collect(); + Value::Array(blocks) + } + }; + converted.push(json!({"role": "user", "content": c})); + } + ChatMessage::Assistant { content, .. } => { + let text = content.clone().unwrap_or_default(); + converted.push(json!({"role": "assistant", "content": text})); + } + ChatMessage::Tool { + tool_call_id, + content, + } => { + // Anthropic 工具结果用 user 角色 + tool_result content block + converted.push(json!({ + "role": "user", + "content": [{ + "type": "tool_result", + "tool_use_id": tool_call_id, + "content": content + }] + })); + } + } + } + + let system = if system_parts.is_empty() { + None + } else { + Some(system_parts.join("\n")) + }; + + (converted, system) + } + + // ==================== 请求体构造 ==================== + + /// 构造 Anthropic chat 请求体 + /// + /// 对齐 Python v1 `AnthropicAdapter.chat` 的 body 构造: + /// - `model` / `messages` / `max_tokens`(必填,默认 1024) + /// - `system`:从 system 消息提取(若存在) + /// - `temperature` / `top_p` / `top_k`:透传 + /// - `stop` → `stop_sequences`(统一为数组) + /// - `reasoning_effort` → `thinking`(对齐 Python ANTHROPIC_MAPPING value_map) + /// - `extra` 透传到顶层 + fn build_chat_body(req: &ChatRequest, stream: bool) -> Value { + let (messages, system) = Self::convert_messages(&req.messages); + let max_tokens = req.max_tokens.unwrap_or(DEFAULT_MAX_TOKENS); + + let mut body = json!({ + "model": req.model, + "messages": messages, + "max_tokens": max_tokens, + }); + + if stream { + body["stream"] = json!(true); + } + if let Some(sys) = system { + body["system"] = json!(sys); + } + if let Some(t) = req.temperature { + body["temperature"] = json!(t); + } + if let Some(top_p) = req.top_p { + body["top_p"] = json!(top_p); + } + if let Some(top_k) = req.top_k { + body["top_k"] = json!(top_k); + } + if let Some(stop) = &req.stop { + let seqs: Vec = match stop { + StopSeq::Single(s) => vec![s.clone()], + StopSeq::Multiple(v) => v.clone(), + }; + body["stop_sequences"] = json!(seqs); + } + + // reasoning_effort → thinking(对齐 Python ANTHROPIC_MAPPING.value_map) + // 仅在未通过 extra 显式设置 thinking 时注入 + if let Some(effort) = req.reasoning_effort { + if body.get("thinking").is_none() { + if let Some(thinking) = thinking_for_effort(effort) { + body["thinking"] = thinking; + } + } + } + + // extra 透传到顶层 + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + + body + } + + // ==================== 响应解析 ==================== + + /// 解析 Anthropic chat 响应 → ChatCompletion + /// + /// Anthropic 响应格式: + /// ```json + /// {"id":"msg_1","model":"claude-...","stop_reason":"end_turn", + /// "content":[{"type":"text","text":"Hello"}], + /// "usage":{"input_tokens":3,"output_tokens":2}} + /// ``` + /// 拼接所有 type==text 的 content block 文本;stop_reason 映射为 finish_reason。 + /// 对应 Python v1 `AnthropicAdapter._parse_response`。 + fn parse_chat_completion(value: &Value, fallback_model: &str) -> Result { + // 拼接所有 text 内容块 + let content: String = value + .get("content") + .and_then(|v| v.as_array()) + .map(|arr| { + arr.iter() + .filter_map(|block| { + if block.get("type").and_then(|t| t.as_str()) == Some("text") { + block + .get("text") + .and_then(|t| t.as_str()) + .map(str::to_owned) + } else { + None + } + }) + .collect::>() + .join("") + }) + .unwrap_or_default(); + + let stop_reason = value.get("stop_reason").and_then(|v| v.as_str()); + let finish_reason = stop_reason.map(map_stop_reason); + + let usage = value.get("usage").and_then(parse_anthropic_usage); + + Ok(ChatCompletion { + id: value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")), + object: "chat.completion".to_string(), + created: value + .get("created") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp), + model: value + .get("model") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| fallback_model.to_string()), + choices: vec![ChatChoice { + index: 0, + message: ChoiceMessage { + role: "assistant".to_string(), + content: Some(content), + tool_calls: None, + }, + finish_reason, + }], + usage, + service_tier: None, + system_fingerprint: None, + }) + } + + /// 解析单个 Anthropic 流式事件 → Option + /// + /// Anthropic 流式事件(JSON 的 `type` 字段): + /// - `content_block_delta`(delta.type==text_delta)→ 增量文本 chunk + /// - `message_delta`(delta.stop_reason)→ 结束原因 chunk(finish_reason) + /// - `message_stop` → 返回 None 并由调用方结束流 + /// - 其他(message_start / content_block_start / content_block_stop / ping)→ None(跳过) + /// + /// 对应 Python v1 `AnthropicAdapter._parse_stream_chunk`。 + fn parse_chunk(value: &Value, fallback_model: &str) -> Option { + let event_type = value.get("type").and_then(|v| v.as_str()).unwrap_or(""); + let chunk_id = value + .get("message") + .and_then(|m| m.get("id")) + .and_then(|v| v.as_str()) + .or_else(|| value.get("id").and_then(|v| v.as_str())) + .map(str::to_owned) + .unwrap_or_else(|| util::generate_id("chatcmpl")); + + match event_type { + "content_block_delta" => { + let delta = value.get("delta").cloned().unwrap_or(Value::Null); + if delta.get("type").and_then(|t| t.as_str()) == Some("text_delta") { + let text = delta + .get("text") + .and_then(|t| t.as_str()) + .unwrap_or("") + .to_string(); + // 空文本不产生 chunk(与 Python 一致:text 为空时不 append) + if text.is_empty() { + return None; + } + Some(make_chunk(chunk_id, fallback_model, Some(text), None)) + } else { + None + } + } + "message_delta" => { + let delta = value.get("delta").cloned().unwrap_or(Value::Null); + let stop_reason = delta.get("stop_reason").and_then(|v| v.as_str()); + stop_reason.map(|r| { + make_chunk( + chunk_id, + fallback_model, + Some(String::new()), + Some(map_stop_reason(r)), + ) + }) + } + // message_stop / message_start / content_block_start / content_block_stop / ping 等 + // 不产生有效 chunk + _ => None, + } + } + + /// 解析 Anthropic /models 响应 → Vec + /// + /// Anthropic 响应:`{"data":[{"id":"...","display_name":"...","type":"model",...}]}`。 + /// `display_name` 作为模型显示名与描述。 + fn parse_models(value: &Value, provider: &str) -> Vec { + let arr = value.get("data").and_then(|v| v.as_array()); + match arr { + Some(arr) => arr + .iter() + .map(|m| { + let id = m + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let display_name = m + .get("display_name") + .and_then(|v| v.as_str()) + .map(str::to_owned); + let model_type = infer_model_type(&id); + ModelInfo { + name: display_name.clone().unwrap_or_else(|| id.clone()), + id, + model_type, + provider: provider.to_string(), + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: matches!(model_type, ModelType::Chat), + description: display_name, + created: None, + } + }) + .collect(), + None => Vec::new(), + } + } +} + +#[async_trait] +impl Adapter for AnthropicAdapter { + fn provider_type(&self) -> &str { + Self::PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + Self::PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + self.compat.capabilities_set().clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 已在构造时初始化,无需额外启动 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // reqwest::Client 自身管理连接池生命周期,无需显式关闭 + Ok(()) + } + + /// 文本对话(Anthropic 特有协议:POST /messages) + async fn chat(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::Chat)?; + let body = Self::build_chat_body(&req, false); + let value = self.post_authed_json("messages", &body).await?; + Self::parse_chat_completion(&value, &req.model) + } + + /// 流式文本对话(Anthropic 特有协议:POST /messages stream=true) + async fn chat_stream(&self, req: ChatRequest) -> Result { + self.ensure_capability(Capabilities::ChatStream)?; + let body = Self::build_chat_body(&req, true); + let url = self.url("messages"); + + let resp = self.send_post(&url, &body).await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + + let model = req.model.clone(); + // 按字节流读取,按行切分解析 SSE(Anthropic 流式同为 SSE data: 格式) + let byte_stream = resp + .bytes_stream() + .map_err(|e| e.to_string()) + .map(|r| r.map(|b| b.to_vec())); + let lines_stream = AnthropicLinesStream::new(byte_stream); + + let stream = async_stream::stream! { + let mut s = lines_stream; + while let Some(line_result) = s.next().await { + let line = match line_result { + Ok(l) => l, + Err(msg) => { + yield Err(AibridgeError::Api { + status: 0, + message: format!("流式读取错误: {msg}"), + }); + return; + } + }; + let line = line.trim(); + // 空行或注释行(心跳)跳过 + if line.is_empty() || line.starts_with(':') { + continue; + } + // 仅处理 data: 行(event: 行由 JSON 内 type 字段判断,跳过) + let data = if let Some(rest) = line.strip_prefix("data: ") { + rest + } else if let Some(rest) = line.strip_prefix("data:") { + rest + } else { + continue; + }; + // 结束标记(Anthropic 通常不发 [DONE],但兼容处理) + if data.trim() == "[DONE]" { + return; + } + match serde_json::from_str::(data) { + Ok(v) => { + // message_stop 事件:结束流 + if v.get("type").and_then(|t| t.as_str()) == Some("message_stop") { + return; + } + match Self::parse_chunk(&v, &model) { + Some(chunk) => yield Ok(chunk), + None => continue, + } + } + // 单行 JSON 解析失败不致命,跳过(与 Python 老版一致) + Err(_) => continue, + } + } + }; + + Ok(stream.boxed()) + } + + /// 模型列表(实时拉取):GET /models + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("models").await?; + let models = Self::parse_models(&value, self.provider_type()); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // image_generate / video_create / video_poll / embed / transcribe / speech / list_voices + // 走 trait 默认实现,返 UnsupportedCapability,与 Python 老版行为一致。 +} + +// ==================== 辅助函数 ==================== + +/// 将统一多模态内容部件转换为 Anthropic content block +/// +/// - Text → `{"type":"text","text":...}` +/// - ImageUrl(data: URI)→ `{"type":"image","source":{"type":"base64",...}}` +/// - ImageUrl(http(s) URL)→ `{"type":"image","source":{"type":"url","url":...}}` +fn convert_content_part(part: &ContentPart) -> Option { + match part { + ContentPart::Text { text } => Some(json!({"type": "text", "text": text})), + ContentPart::ImageUrl { image_url } => { + let url = &image_url.url; + if let Some(rest) = url.strip_prefix("data:") { + // data:;, + let (media_type, data) = match rest.split_once(',') { + Some((meta, d)) => { + let mt = meta.split(';').next().unwrap_or("image/png"); + (mt, d) + } + None => return None, + }; + Some(json!({ + "type": "image", + "source": {"type": "base64", "media_type": media_type, "data": data} + })) + } else { + // http(s) URL:Anthropic 新版支持 url source + Some(json!({ + "type": "image", + "source": {"type": "url", "url": url} + })) + } + } + } +} + +/// 将 Anthropic stop_reason 映射为统一 finish_reason +/// +/// 对应 Python v1 `stop_reason_map`: +/// - end_turn → stop +/// - max_tokens → length +/// - stop_sequence → stop +/// - tool_use → tool_calls +/// - 其他 → 原值透传 +fn map_stop_reason(reason: &str) -> String { + match reason { + "end_turn" => "stop".into(), + "max_tokens" => "length".into(), + "stop_sequence" => "stop".into(), + "tool_use" => "tool_calls".into(), + other => other.into(), + } +} + +/// 将 reasoning_effort 映射为 Anthropic thinking 配置 +/// +/// 对齐 Python v1 `ANTHROPIC_MAPPING.value_map["reasoning_effort"]`: +/// - Low → `{"type":"enabled","budget_tokens":1024}` +/// - Medium → `{"type":"enabled","budget_tokens":4096}` +/// - High → `{"type":"enabled","budget_tokens":16384}` +/// - None / Auto → 不注入(返回 None) +/// +/// 注意:Anthropic 要求 `budget_tokens < max_tokens`,调用方需确保 max_tokens 足够大。 +fn thinking_for_effort(effort: ReasoningEffort) -> Option { + match effort { + ReasoningEffort::Low => Some(json!({"type": "enabled", "budget_tokens": 1024})), + ReasoningEffort::Medium => Some(json!({"type": "enabled", "budget_tokens": 4096})), + ReasoningEffort::High => Some(json!({"type": "enabled", "budget_tokens": 16384})), + ReasoningEffort::None | ReasoningEffort::Auto => None, + } +} + +/// 构造单个流式 chunk(简化重复代码) +fn make_chunk( + id: String, + model: &str, + content: Option, + finish_reason: Option, +) -> ChatCompletionChunk { + ChatCompletionChunk { + id, + object: "chat.completion.chunk".to_string(), + created: util::current_timestamp(), + model: model.to_string(), + choices: vec![ChatCompletionDelta { + index: 0, + delta: DeltaMessage { + role: Some("assistant".to_string()), + content, + tool_calls: None, + }, + finish_reason, + }], + usage: None, + } +} + +/// 解析 Anthropic usage 统计 +/// +/// Anthropic usage 格式:`{"input_tokens": N, "output_tokens": M}`, +/// 需转换为统一 ChatUsage(prompt / completion / total)。 +fn parse_anthropic_usage(v: &Value) -> Option { + let prompt = v.get("input_tokens").and_then(|x| x.as_u64())?; + let completion = v.get("output_tokens").and_then(|x| x.as_u64()).unwrap_or(0); + Some(crate::model::chat::ChatUsage { + prompt_tokens: prompt, + completion_tokens: completion, + total_tokens: prompt + completion, + }) +} + +// ==================== SSE 行流适配器 ==================== + +/// 将字节流按行切分的适配器(Anthropic 流式用) +/// +/// 与 `openai_compat::LinesStream` 等价,独立实现避免引用其私有结构。 +struct AnthropicLinesStream { + inner: S, + buffer: Vec, +} + +impl AnthropicLinesStream { + fn new(inner: S) -> Self { + Self { + inner, + buffer: Vec::new(), + } + } +} + +impl futures::Stream for AnthropicLinesStream +where + S: futures::Stream, String>> + Unpin, +{ + type Item = std::result::Result; + + fn poll_next( + mut self: std::pin::Pin<&mut Self>, + cx: &mut std::task::Context<'_>, + ) -> std::task::Poll> { + use std::task::Poll; + loop { + // 先看缓冲区是否已有完整行 + if let Some(pos) = self.buffer.iter().position(|&b| b == b'\n') { + let mut line: Vec = self.buffer.drain(..=pos).collect(); + // 去掉末尾 \n 与可能的 \r + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + // 缓冲区无完整行,拉取下一 chunk + match std::pin::Pin::new(&mut self.inner).poll_next(cx) { + Poll::Ready(Some(Err(msg))) => return Poll::Ready(Some(Err(msg))), + Poll::Ready(Some(Ok(chunk))) => { + self.buffer.extend_from_slice(&chunk); + // 继续循环,尝试从缓冲区切出行 + } + Poll::Ready(None) => { + // 流结束,把缓冲区剩余内容作为最后一行返回 + if !self.buffer.is_empty() { + let mut line = std::mem::take(&mut self.buffer); + if line.last() == Some(&b'\n') { + line.pop(); + } + if line.last() == Some(&b'\r') { + line.pop(); + } + let s = String::from_utf8_lossy(&line).into_owned(); + return Poll::Ready(Some(Ok(s))); + } + return Poll::Ready(None); + } + Poll::Pending => return Poll::Pending, + } + } + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::http::HttpClient; + use crate::model::image::ImageRequest; + use crate::model::options::{EmbedInput, EmbedRequest, ReasoningEffort}; + use crate::model::video::VideoRequest; + use mockito::Server; + use std::collections::HashMap; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 AnthropicAdapter(指向 mockito server) + fn make_anthropic(server: &Server) -> AnthropicAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("anthropic", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let compat = OpenAiCompatAdapter::with_http( + http, + config, + "anthropic", + "Anthropic Claude", + anthropic_capabilities(), + ); + AnthropicAdapter::with_compat(compat) + } + + /// 标准 Anthropic chat 成功响应体 + fn anthropic_chat_body() -> Value { + json!({ + "id": "msg_1", + "type": "message", + "role": "assistant", + "model": "claude-3-5-sonnet-20241022", + "content": [{"type": "text", "text": "Hello!"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 5} + }) + } + + // ==================== 元信息测试 ==================== + + #[tokio::test] + async fn provider_type_and_name() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + assert_eq!(adapter.provider_type(), "anthropic"); + assert_eq!(adapter.provider_name(), "Anthropic Claude"); + } + + #[tokio::test] + async fn requires_api_key() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + assert!(adapter.requires_api_key()); + } + + #[tokio::test] + async fn capabilities_include_chat_and_stream() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::Chat)); + assert!(caps.contains(&Capabilities::ChatStream)); + assert!(caps.contains(&Capabilities::Vision)); + // 不支持 image / embed / video + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::Embedding)); + assert!(!caps.contains(&Capabilities::VideoGenerate)); + } + + #[tokio::test] + async fn start_close_are_noops() { + let server = Server::new_async().await; + let mut adapter = make_anthropic(&server); + adapter.start().await.unwrap(); + adapter.close().await.unwrap(); + } + + #[test] + fn default_base_url_matches_spec() { + assert_eq!(DEFAULT_ANTHROPIC_BASE_URL, "https://api.anthropic.com/v1"); + assert_eq!(ANTHROPIC_API_VERSION, "2023-06-01"); + } + + // ==================== chat 正常路径 ==================== + + #[tokio::test] + async fn chat_success_parses_completion() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_header("x-api-key", "test-key") + .match_header("anthropic-version", "2023-06-01") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.expect("chat 应成功"); + + assert_eq!(resp.id, "msg_1"); + assert_eq!(resp.model, "claude-3-5-sonnet-20241022"); + assert_eq!(resp.choices.len(), 1); + assert_eq!(resp.choices[0].message.content.as_deref(), Some("Hello!")); + assert_eq!(resp.choices[0].finish_reason.as_deref(), Some("stop")); + // usage 解析:input/output → prompt/completion/total + let usage = resp.usage.as_ref().expect("应有 usage"); + assert_eq!(usage.prompt_tokens, 10); + assert_eq!(usage.completion_tokens, 5); + assert_eq!(usage.total_tokens, 15); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_sends_x_api_key_and_anthropic_version_headers() { + // 验证 Anthropic 特有认证 header(x-api-key + anthropic-version),非 Bearer + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_header("x-api-key", "test-key") + .match_header("anthropic-version", "2023-06-01") + .match_header("authorization", mockito::Matcher::Missing) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_extracts_system_to_top_level_field() { + // system 消息应提取到顶层 system 字段,不在 messages 中 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "claude-3-5-sonnet-20241022", + "system": "you are helpful", + "messages": [{"role": "user", "content": "hi"}], + "max_tokens": 100 + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder( + "claude-3-5-sonnet-20241022", + vec![ + ChatMessage::system("you are helpful"), + ChatMessage::user("hi"), + ], + ) + .max_tokens(100) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_converts_user_and_assistant_messages() { + // 多轮对话:user/assistant 消息转为 Anthropic messages(system 提取到顶层) + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "messages": [ + {"role": "user", "content": "first question"}, + {"role": "assistant", "content": "first answer"}, + {"role": "user", "content": "second question"} + ] + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder( + "claude-3-5-sonnet-20241022", + vec![ + ChatMessage::user("first question"), + ChatMessage::assistant("first answer"), + ChatMessage::user("second question"), + ], + ) + .max_tokens(100) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_default_max_tokens_when_not_provided() { + // max_tokens 必填,未指定时兜底 1024 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "max_tokens": 1024 + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_maps_stop_to_stop_sequences() { + // stop → stop_sequences(统一为数组) + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "stop_sequences": ["END", "STOP"] + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .stop(StopSeq::Multiple(vec!["END".into(), "STOP".into()])) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_passes_temperature_top_p_top_k() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "temperature": 0.7, + "top_p": 0.9, + "top_k": 40 + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .temperature(0.7) + .top_p(0.9) + .top_k(40) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_passes_extra_params_through() { + // extra 透传到顶层 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "custom_param": "custom_value" + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .extra("custom_param", "custom_value") + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_injects_thinking_for_reasoning_effort() { + // reasoning_effort → thinking(对齐 Python ANTHROPIC_MAPPING.value_map) + // Anthropic 协议不接受 reasoning_effort 字段,仅注入 thinking + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "thinking": {"type": "enabled", "budget_tokens": 16384}, + "max_tokens": 20000 + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(20000) + .reasoning_effort(ReasoningEffort::High) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_no_thinking_without_reasoning_effort() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "claude-3-5-sonnet-20241022", + "messages": [{"role": "user", "content": "hi"}] + }))) + .with_status(200) + .with_body(anthropic_chat_body().to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let _ = adapter.chat(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_concatenates_multiple_text_blocks() { + // content 数组含多个 text 块时,应拼接为完整文本 + let mut server = Server::new_async().await; + let body = json!({ + "id": "msg_2", + "model": "claude-3-5-sonnet-20241022", + "content": [ + {"type": "text", "text": "Hello"}, + {"type": "text", "text": " world"} + ], + "stop_reason": "end_turn", + "usage": {"input_tokens": 5, "output_tokens": 2} + }); + server + .mock("POST", "/messages") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!( + resp.choices[0].message.content.as_deref(), + Some("Hello world") + ); + } + + #[tokio::test] + async fn chat_maps_various_stop_reasons() { + let cases = [ + ("end_turn", "stop"), + ("max_tokens", "length"), + ("stop_sequence", "stop"), + ("tool_use", "tool_calls"), + ]; + for (stop_reason, expected_finish) in cases { + let mut server = Server::new_async().await; + let body = json!({ + "id": "msg_x", + "model": "claude-3-5-sonnet-20241022", + "content": [{"type": "text", "text": "ok"}], + "stop_reason": stop_reason, + "usage": {"input_tokens": 1, "output_tokens": 1} + }); + server + .mock("POST", "/messages") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + let adapter = make_anthropic(&server); + let req = + ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let resp = adapter.chat(req).await.unwrap(); + assert_eq!( + resp.choices[0].finish_reason.as_deref(), + Some(expected_finish), + "stop_reason={stop_reason} 应映射为 {expected_finish}" + ); + } + } + + // ==================== chat 错误路径 ==================== + + #[tokio::test] + async fn chat_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(401) + .with_body( + json!({ + "type": "error", + "error": {"type": "authentication_error", "message": "invalid x-api-key"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn chat_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(429) + .with_body( + json!({ + "type": "error", + "error": {"type": "rate_limit_error", "message": "rate limit exceeded"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn chat_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(404) + .with_body( + json!({ + "type": "error", + "error": {"type": "not_found_error", "message": "model not found"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("bad-model", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn chat_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(400) + .with_body( + json!({ + "type": "error", + "error": {"type": "invalid_request_error", "message": "max_tokens is invalid"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("max_tokens")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn chat_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(500) + .with_body( + json!({ + "type": "error", + "error": {"type": "api_error", "message": "internal server error"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn chat_unsupported_capability_returns_error() { + // 不支持 Chat 能力时应返 UnsupportedCapability + let server = Server::new_async().await; + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .build(); + let config = ProviderConfig::from_options("anthropic", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + let compat = OpenAiCompatAdapter::with_http( + http, + config, + "anthropic", + "Anthropic Claude", + CapabilitySet::new(), + ); + let adapter = AnthropicAdapter::with_compat(compat); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== chat_stream 正常 + 错误路径 ==================== + + #[tokio::test] + async fn chat_stream_parses_sse_events() { + // Anthropic 流式:content_block_delta + message_delta + message_stop + let mut server = Server::new_async().await; + let sse = "event: message_start\n\ + data: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\",\"model\":\"claude-3-5-sonnet-20241022\"}}\n\n\ + event: content_block_delta\n\ + data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hello\"}}\n\n\ + event: content_block_delta\n\ + data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\" world\"}}\n\n\ + event: message_delta\n\ + data: {\"type\":\"message_delta\",\"delta\":{\"stop_reason\":\"end_turn\"}}\n\n\ + event: message_stop\n\ + data: {\"type\":\"message_stop\"}\n\n"; + server + .mock("POST", "/messages") + .with_status(200) + .with_header("content-type", "text/event-stream") + .with_body(sse) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let mut stream = adapter.chat_stream(req).await.expect("stream 应建立"); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + // 2 个文本 delta + 1 个 message_delta(finish_reason) + assert_eq!(chunks.len(), 3); + // 拼接前两块内容 + let mut content = String::new(); + content.push_str(chunks[0].choices[0].delta.content.as_deref().unwrap_or("")); + content.push_str(chunks[1].choices[0].delta.content.as_deref().unwrap_or("")); + assert_eq!(content, "Hello world"); + // 第三块是 message_delta,finish_reason = stop + assert_eq!(chunks[2].choices[0].finish_reason.as_deref(), Some("stop")); + } + + #[tokio::test] + async fn chat_stream_skips_non_text_delta_events() { + // content_block_start / ping 等事件应被跳过 + let mut server = Server::new_async().await; + let sse = "event: message_start\n\ + data: {\"type\":\"message_start\",\"message\":{\"id\":\"msg_1\"}}\n\n\ + event: content_block_start\n\ + data: {\"type\":\"content_block_start\",\"index\":0,\"content_block\":{\"type\":\"text\",\"text\":\"\"}}\n\n\ + event: content_block_delta\n\ + data: {\"type\":\"content_block_delta\",\"index\":0,\"delta\":{\"type\":\"text_delta\",\"text\":\"Hi\"}}\n\n\ + event: ping\n\ + data: {\"type\":\"ping\"}\n\n\ + event: message_stop\n\ + data: {\"type\":\"message_stop\"}\n\n"; + server + .mock("POST", "/messages") + .with_status(200) + .with_body(sse) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + let mut chunks = Vec::new(); + while let Some(chunk) = stream.next().await { + chunks.push(chunk.unwrap()); + } + // 仅 1 个有效 chunk(content_block_delta) + assert_eq!(chunks.len(), 1); + assert_eq!(chunks[0].choices[0].delta.content.as_deref(), Some("Hi")); + } + + #[tokio::test] + async fn chat_stream_sends_stream_true_and_headers() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/messages") + .match_header("x-api-key", "test-key") + .match_header("anthropic-version", "2023-06-01") + .match_body(mockito::Matcher::PartialJson(json!({ + "stream": true, + "max_tokens": 100 + }))) + .with_status(200) + .with_body("event: message_stop\ndata: {\"type\":\"message_stop\"}\n\n") + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let mut stream = adapter.chat_stream(req).await.unwrap(); + while stream.next().await.is_some() {} + mock.assert_async().await; + } + + #[tokio::test] + async fn chat_stream_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(401) + .with_body( + json!({"type":"error","error":{"type":"authentication_error","message":"bad key"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::Authentication { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + #[tokio::test] + async fn chat_stream_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/messages") + .with_status(429) + .with_body( + json!({"type":"error","error":{"type":"rate_limit_error","message":"slow down"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(100) + .build(); + let result = adapter.chat_stream(req).await; + match result { + Err(e) => assert!(matches!(e, AibridgeError::RateLimit { .. })), + Ok(_) => panic!("chat_stream 应返回错误而非 stream"), + } + } + + // ==================== list_models ==================== + + #[tokio::test] + async fn list_models_success() { + let mut server = Server::new_async().await; + let body = json!({ + "data": [ + {"id": "claude-3-5-sonnet-20241022", "display_name": "Claude 3.5 Sonnet", "type": "model"}, + {"id": "claude-3-opus-20240229", "display_name": "Claude 3 Opus", "type": "model"} + ] + }); + server + .mock("GET", "/models") + .match_header("x-api-key", "test-key") + .match_header("anthropic-version", "2023-06-01") + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "claude-3-5-sonnet-20241022"); + assert_eq!(models[0].name, "Claude 3.5 Sonnet"); + assert_eq!(models[0].provider, "anthropic"); + assert_eq!(models[0].description.as_deref(), Some("Claude 3.5 Sonnet")); + } + + #[tokio::test] + async fn list_models_error_401() { + let mut server = Server::new_async().await; + server + .mock("GET", "/models") + .with_status(401) + .with_body( + json!({"type":"error","error":{"type":"authentication_error","message":"bad key"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_anthropic(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ==================== 不支持能力 ==================== + + #[tokio::test] + async fn image_generate_unsupported() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + let req = ImageRequest::builder("claude-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn video_create_unsupported() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + let req = VideoRequest::builder("claude-3", "a cat walking").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn embed_unsupported() { + let server = Server::new_async().await; + let adapter = make_anthropic(&server); + let req = EmbedRequest { + model: "claude-3".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ==================== 辅助函数单元测试 ==================== + + #[test] + fn anthropic_mapping_renames_stop_to_stop_sequences() { + let pm = anthropic_mapping(); + let mut params = HashMap::new(); + params.insert("stop".to_string(), json!("END")); + params.insert("max_tokens".to_string(), json!(1000)); + params.insert("temperature".to_string(), json!(0.7)); + let result = pm.apply(¶ms); + // stop → stop_sequences + assert_eq!( + result.get("stop_sequences").and_then(|v| v.as_str()), + Some("END") + ); + assert!(!result.contains_key("stop")); + // 其余原名透传 + assert_eq!( + result.get("max_tokens").and_then(|v| v.as_i64()), + Some(1000) + ); + assert_eq!( + result.get("temperature").and_then(|v| v.as_f64()), + Some(0.7) + ); + } + + #[test] + fn map_stop_reason_matches_python() { + assert_eq!(map_stop_reason("end_turn"), "stop"); + assert_eq!(map_stop_reason("max_tokens"), "length"); + assert_eq!(map_stop_reason("stop_sequence"), "stop"); + assert_eq!(map_stop_reason("tool_use"), "tool_calls"); + // 未知值原样透传 + assert_eq!(map_stop_reason("other"), "other"); + } + + #[test] + fn thinking_for_effort_maps_levels() { + assert_eq!( + thinking_for_effort(ReasoningEffort::Low), + Some(json!({"type": "enabled", "budget_tokens": 1024})) + ); + assert_eq!( + thinking_for_effort(ReasoningEffort::Medium), + Some(json!({"type": "enabled", "budget_tokens": 4096})) + ); + assert_eq!( + thinking_for_effort(ReasoningEffort::High), + Some(json!({"type": "enabled", "budget_tokens": 16384})) + ); + assert_eq!(thinking_for_effort(ReasoningEffort::None), None); + assert_eq!(thinking_for_effort(ReasoningEffort::Auto), None); + } + + #[test] + fn convert_messages_extracts_system_and_converts_roles() { + let msgs = vec![ + ChatMessage::system("be helpful"), + ChatMessage::user("q1"), + ChatMessage::assistant("a1"), + ChatMessage::user("q2"), + ]; + let (converted, system) = AnthropicAdapter::convert_messages(&msgs); + assert_eq!(system.as_deref(), Some("be helpful")); + // system 不在 messages 中 + assert_eq!(converted.len(), 3); + assert_eq!(converted[0]["role"], "user"); + assert_eq!(converted[0]["content"], "q1"); + assert_eq!(converted[1]["role"], "assistant"); + assert_eq!(converted[1]["content"], "a1"); + assert_eq!(converted[2]["role"], "user"); + assert_eq!(converted[2]["content"], "q2"); + } + + #[test] + fn convert_messages_joins_multiple_system_messages() { + let msgs = vec![ + ChatMessage::system("rule 1"), + ChatMessage::system("rule 2"), + ChatMessage::user("hi"), + ]; + let (converted, system) = AnthropicAdapter::convert_messages(&msgs); + // 多条 system 用 \n 拼接 + assert_eq!(system.as_deref(), Some("rule 1\nrule 2")); + assert_eq!(converted.len(), 1); + } + + #[test] + fn convert_messages_multimodal_user_to_content_blocks() { + let msgs = vec![ChatMessage::user_multimodal(vec![ + ContentPart::Text { + text: "describe this".into(), + }, + ContentPart::ImageUrl { + image_url: crate::model::chat::ImageUrl::new("data:image/png;base64,aGVsbG8="), + }, + ])]; + let (converted, system) = AnthropicAdapter::convert_messages(&msgs); + assert!(system.is_none()); + assert_eq!(converted.len(), 1); + assert_eq!(converted[0]["role"], "user"); + let content = converted[0]["content"].as_array().expect("应为数组"); + assert_eq!(content.len(), 2); + // 文本块 + assert_eq!(content[0]["type"], "text"); + assert_eq!(content[0]["text"], "describe this"); + // 图片块(data: URI → base64 source) + assert_eq!(content[1]["type"], "image"); + assert_eq!(content[1]["source"]["type"], "base64"); + assert_eq!(content[1]["source"]["media_type"], "image/png"); + assert_eq!(content[1]["source"]["data"], "aGVsbG8="); + } + + #[test] + fn convert_content_part_url_image_uses_url_source() { + let part = ContentPart::ImageUrl { + image_url: crate::model::chat::ImageUrl::new("https://example.com/x.png"), + }; + let block = convert_content_part(&part).expect("应返回 Some"); + assert_eq!(block["type"], "image"); + assert_eq!(block["source"]["type"], "url"); + assert_eq!(block["source"]["url"], "https://example.com/x.png"); + } + + #[test] + fn parse_chat_completion_extracts_text_and_usage() { + let v = json!({ + "id": "msg_9", + "model": "claude-3-opus-20240229", + "content": [{"type": "text", "text": "response text"}], + "stop_reason": "end_turn", + "usage": {"input_tokens": 10, "output_tokens": 20} + }); + let cc = AnthropicAdapter::parse_chat_completion(&v, "fallback").unwrap(); + assert_eq!(cc.id, "msg_9"); + assert_eq!(cc.model, "claude-3-opus-20240229"); + assert_eq!( + cc.choices[0].message.content.as_deref(), + Some("response text") + ); + assert_eq!(cc.choices[0].finish_reason.as_deref(), Some("stop")); + let usage = cc.usage.unwrap(); + assert_eq!(usage.prompt_tokens, 10); + assert_eq!(usage.completion_tokens, 20); + assert_eq!(usage.total_tokens, 30); + } + + #[test] + fn parse_chat_completion_without_usage() { + let v = json!({ + "id": "msg_x", + "content": [{"type": "text", "text": "no usage"}], + "stop_reason": "max_tokens" + }); + let cc = AnthropicAdapter::parse_chat_completion(&v, "fallback").unwrap(); + assert_eq!(cc.choices[0].message.content.as_deref(), Some("no usage")); + assert_eq!(cc.choices[0].finish_reason.as_deref(), Some("length")); + assert!(cc.usage.is_none()); + // model 回退到 fallback + assert_eq!(cc.model, "fallback"); + } + + #[test] + fn parse_chunk_content_block_delta_returns_text() { + let v = json!({ + "type": "content_block_delta", + "delta": {"type": "text_delta", "text": "hello"} + }); + let chunk = AnthropicAdapter::parse_chunk(&v, "claude-3").expect("应返回 Some"); + assert_eq!(chunk.choices[0].delta.content.as_deref(), Some("hello")); + assert!(chunk.choices[0].finish_reason.is_none()); + } + + #[test] + fn parse_chunk_empty_text_returns_none() { + let v = json!({ + "type": "content_block_delta", + "delta": {"type": "text_delta", "text": ""} + }); + assert!(AnthropicAdapter::parse_chunk(&v, "claude-3").is_none()); + } + + #[test] + fn parse_chunk_message_delta_returns_finish_reason() { + let v = json!({ + "type": "message_delta", + "delta": {"stop_reason": "end_turn"} + }); + let chunk = AnthropicAdapter::parse_chunk(&v, "claude-3").expect("应返回 Some"); + assert_eq!(chunk.choices[0].finish_reason.as_deref(), Some("stop")); + assert_eq!(chunk.choices[0].delta.content.as_deref(), Some("")); + } + + #[test] + fn parse_chunk_skips_non_text_delta_types() { + // 非 text_delta 的 delta(如 input_json_delta)应跳过 + let v = json!({ + "type": "content_block_delta", + "delta": {"type": "input_json_delta", "partial_json": "{}"} + }); + assert!(AnthropicAdapter::parse_chunk(&v, "claude-3").is_none()); + } + + #[test] + fn parse_chunk_skips_message_start_and_ping() { + assert!( + AnthropicAdapter::parse_chunk(&json!({"type":"message_start"}), "claude-3").is_none() + ); + assert!(AnthropicAdapter::parse_chunk(&json!({"type":"ping"}), "claude-3").is_none()); + assert!( + AnthropicAdapter::parse_chunk(&json!({"type":"content_block_start"}), "claude-3") + .is_none() + ); + } + + #[test] + fn parse_models_uses_id_and_display_name() { + let v = json!({ + "data": [ + {"id": "claude-3-5-sonnet-20241022", "display_name": "Claude 3.5 Sonnet"}, + {"id": "claude-3-haiku-20240307", "display_name": "Claude 3 Haiku"} + ] + }); + let models = AnthropicAdapter::parse_models(&v, "anthropic"); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "claude-3-5-sonnet-20241022"); + assert_eq!(models[0].name, "Claude 3.5 Sonnet"); + assert_eq!(models[0].provider, "anthropic"); + assert_eq!(models[1].id, "claude-3-haiku-20240307"); + } + + #[test] + fn parse_models_empty_returns_vec() { + assert!(AnthropicAdapter::parse_models(&json!({"data": []}), "anthropic").is_empty()); + assert!(AnthropicAdapter::parse_models(&json!({}), "anthropic").is_empty()); + } + + #[test] + fn parse_anthropic_usage_maps_tokens() { + let v = json!({"input_tokens": 7, "output_tokens": 3}); + let usage = parse_anthropic_usage(&v).expect("应返回 Some"); + assert_eq!(usage.prompt_tokens, 7); + assert_eq!(usage.completion_tokens, 3); + assert_eq!(usage.total_tokens, 10); + } + + #[test] + fn parse_anthropic_usage_missing_input_returns_none() { + assert!(parse_anthropic_usage(&json!({"output_tokens": 3})).is_none()); + } + + #[test] + fn build_chat_body_includes_required_fields() { + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(50) + .temperature(0.5) + .build(); + let body = AnthropicAdapter::build_chat_body(&req, false); + assert_eq!(body["model"], "claude-3-5-sonnet-20241022"); + assert_eq!(body["max_tokens"], 50); + assert_eq!(body["temperature"], 0.5); + assert!(body.get("messages").is_some()); + // 非 stream 模式不应有 stream 字段 + assert!(body.get("stream").is_none()); + } + + #[test] + fn build_chat_body_stream_adds_stream_true() { + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(50) + .build(); + let body = AnthropicAdapter::build_chat_body(&req, true); + assert_eq!(body["stream"], true); + } + + #[test] + fn build_chat_body_single_stop_becomes_array() { + let req = ChatRequest::builder("claude-3-5-sonnet-20241022", vec![ChatMessage::user("hi")]) + .max_tokens(50) + .stop(StopSeq::Single("END".into())) + .build(); + let body = AnthropicAdapter::build_chat_body(&req, false); + assert_eq!(body["stop_sequences"], json!(["END"])); + } +} From e6151da072e2376e02fd26bfcc4c71d71f5705b4 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 00:39:59 +0800 Subject: [PATCH 34/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20stability=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../aibridge-core/src/adapters/stability.rs | 1450 +++++++++++++++++ 1 file changed, 1450 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/stability.rs diff --git a/crates/aibridge-core/src/adapters/stability.rs b/crates/aibridge-core/src/adapters/stability.rs new file mode 100644 index 0000000..49cdd85 --- /dev/null +++ b/crates/aibridge-core/src/adapters/stability.rs @@ -0,0 +1,1450 @@ +//! Stability AI 适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/stability.py`。 +//! +//! Stability AI 为**独立协议**(非 OpenAI 兼容),不复用 `OpenAiCompatAdapter`: +//! - 图像生成:`POST /v1/generation/{engine_id}/text-to-image`(同步,返回 base64 图像) +//! - 模型列表:`GET /v1/engines/list`(返回顶层数组) +//! - 认证:`Authorization: Bearer ` +//! +//! ## 协议要点 +//! - 请求体用 `text_prompts` 数组(每项含 `text` + `weight`,正向 weight=1,负向 weight=-1) +//! - 尺寸用 `width` / `height` 数值字段(非 size 字符串) +//! - 采样参数:`steps` / `seed` / `cfg_scale` / `samples`(生成数量,上限 10)/ `style_preset` +//! - 响应:`{"artifacts": [{"base64": "..."}]}`(注意是 `base64` 而非 `b64_json`) +//! +//! ## 能力范围 +//! 仅 `ImageGenerate`。chat / video / embed / audio 走 `Adapter` trait 默认实现 +//! 返 `UnsupportedCapability`,与 Python v1 抛 `UnsupportedCapabilityError` 行为一致。 +//! +//! ## 错误映射(对齐任务要求,与 Python v1 略有差异) +//! - 401 → Authentication +//! - 429 → RateLimit +//! - 400 → Validation(Python v1 走 Api,此处按统一错误规范映射为 Validation) +//! - 5xx 及其余 ≥400 → Api(提取 message) + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::common::{infer_model_type, ModelInfo, ModelType}; +use crate::model::image::{ImageData, ImageRequest, ImageResult}; +use crate::util; + +// ==================== 默认配置 ==================== + +/// Stability AI 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`。 +pub const DEFAULT_STABILITY_BASE_URL: &str = "https://api.stability.ai"; + +/// Stability AI 默认引擎(SDXL 1.1) +/// +/// 对应 Python v1 `DEFAULT_ENGINE`。 +pub const DEFAULT_ENGINE: &str = "stable-diffusion-xl-1024-v1-1"; + +/// 默认图像尺寸(与 Python v1 一致) +const DEFAULT_WIDTH: u32 = 1024; +const DEFAULT_HEIGHT: u32 = 1024; + +/// samples(生成数量)上限(与 Python v1 `min(samples, 10)` 一致) +const MAX_SAMPLES: u32 = 10; + +// ==================== 能力集合 ==================== + +/// Stability 支持的能力集合 +/// +/// 对齐 Python v1 `StabilityAdapter.supported_capabilities = ["image"]`。 +/// Rust 用 `ImageGenerate` 表达图像生成能力。 +fn stability_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::ImageGenerate); + caps +} + +// ==================== StabilityAdapter ==================== + +/// Stability AI 适配器 +/// +/// 持有 HTTP 客户端与 Provider 配置,实现 Stability 独立图像生成协议。 +/// 不复用 `OpenAiCompatAdapter`(请求体结构、响应字段与 OpenAI 不一致)。 +/// +/// ## 阶段范围 +/// 阶段 2b 实现 `image_generate`(文生图)+ `list_models`(实时拉取 engines)。 +/// 图像编辑(image-to-image,multipart)待后续阶段补齐 `image_edit` trait 方法后再实现。 +pub struct StabilityAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl StabilityAdapter { + /// 创建 Stability AI 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_STABILITY_BASE_URL`] 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_STABILITY_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: stability_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_STABILITY_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: stability_capabilities(), + } + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: stability)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用 Stability 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .json(body) + .send() + .await + .map_err(map_reqwest_error)?; + Self::parse_json_response(resp).await + } + + /// 发送带 Bearer 认证的 GET 请求,并用 Stability 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .send() + .await + .map_err(map_reqwest_error)?; + Self::parse_json_response(resp).await + } + + /// 统一解析响应:状态码非 2xx 走错误映射,成功则解析为 JSON + async fn parse_json_response(resp: reqwest::Response) -> Result { + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + // ==================== 内部:请求体构造 ==================== + + /// 构造 Stability text-to-image 请求体 + /// + /// 移植自 Python v1 `image_generate`: + /// - `text_prompts`:正向 weight=1,负向 weight=-1 + /// - `width` / `height`:从 width/height/size 解析,默认 1024x1024 + /// - `samples`:生成数量,上限 10 + /// - `steps` / `seed` / `cfg_scale` / `style_preset`:透传 + /// - `extra`:合并到顶层(透传厂商特有参数) + fn build_generate_body(&self, req: &ImageRequest) -> Value { + // text_prompts:正向提示词 + let mut text_prompts = vec![json!({ "text": req.prompt, "weight": 1 })]; + // 负面提示词 weight=-1 + if let Some(np) = &req.negative_prompt { + text_prompts.push(json!({ "text": np, "weight": -1 })); + } + + // 尺寸解析 + let (width, height) = resolve_dimensions(req.width, req.height, req.size.as_deref()); + + let mut body = json!({ + "text_prompts": text_prompts, + "width": width, + "height": height, + }); + + // 采样参数(可选) + if let Some(steps) = req.steps { + body["steps"] = json!(steps); + } + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + if let Some(cfg) = req.cfg_scale { + body["cfg_scale"] = json!(cfg); + } + // samples 上限 10 + body["samples"] = json!(req.n.min(MAX_SAMPLES)); + // 风格预设 + if let Some(style) = &req.style { + body["style_preset"] = json!(style); + } + + // extra 透传(合并到顶层) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + // ==================== 内部:响应解析 ==================== + + /// 解析 Stability text-to-image 响应 → ImageResult + /// + /// Stability 响应结构: + /// ```json + /// { "artifacts": [{"base64": "...", "finishReason": "SUCCESS", "seed": 123}] } + /// ``` + /// 注意:Stability 用 `base64` 而非 `b64_json`,无 `url` 字段。 + /// artifacts 为空时返 Api 错误(与 Python v1 `raise APIError("No image generated")` 一致)。 + fn parse_image_result(value: &Value, fallback_model: &str) -> Result { + let data: Vec = value + .get("artifacts") + .and_then(|v| v.as_array()) + .map(|arr| arr.iter().filter_map(parse_artifact).collect()) + .unwrap_or_default(); + + if data.is_empty() { + return Err(AibridgeError::Api { + status: 0, + message: "Stability 未返回图像数据 (artifacts 为空)".to_string(), + }); + } + + Ok(ImageResult { + id: util::generate_id("img"), + object: "image.generation".to_string(), + created: util::current_timestamp(), + model: fallback_model.to_string(), + data, + }) + } + + /// 解析 Stability `/v1/engines/list` 响应 → Vec + /// + /// Stability engines/list 返回**顶层数组**(非 `{"data": [...]}`),格式: + /// ```json + /// [{"id": "stable-diffusion-xl-1024-v1-1", "name": "...", "description": "...", "create_time": 123}] + /// ``` + /// 兼容 `{"data": [...]}` 包装格式。模型类型由 `infer_model_type` 推断 + /// (`stable-diffusion` 关键字 → Image)。 + fn parse_engines(value: &Value) -> Vec { + let arr: Vec<&Value> = match value { + Value::Array(a) => a.iter().collect(), + Value::Object(_) => value + .get("data") + .and_then(|v| v.as_array()) + .map(|a| a.iter().collect()) + .unwrap_or_default(), + _ => Vec::new(), + }; + + arr.iter() + .map(|m| { + let id = m + .get("id") + .and_then(|v| v.as_str()) + .unwrap_or("") + .to_string(); + let name = m + .get("name") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .unwrap_or_else(|| id.clone()); + ModelInfo { + model_type: infer_model_type(&id), + provider: "stability".to_string(), + description: m + .get("description") + .and_then(|v| v.as_str()) + .map(str::to_owned), + created: m.get("create_time").and_then(|v| v.as_u64()), + id, + name, + capabilities: Vec::new(), + max_tokens: None, + supports_streaming: false, + } + }) + .collect() + } + + // ==================== 内部:错误映射 ==================== + + /// 将 Stability API 错误响应映射为 AibridgeError + /// + /// 对齐任务要求(与 Python v1 `_handle_stability_error` 略有差异): + /// - 401 → Authentication("Invalid Stability API key") + /// - 429 → RateLimit("Stability rate limit exceeded") + /// - 400 → Validation(携带 details,Python v1 走 Api) + /// - 其余 ≥400 → Api(提取 message,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Stability API key".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Stability rate limit exceeded".to_string(), + retry_after: None, + }, + 400 => { + let message = parse_error_message(body, status); + let details = + serde_json::from_str::(body).unwrap_or(serde_json::Value::Null); + AibridgeError::Validation { message, details } + } + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for StabilityAdapter { + fn provider_type(&self) -> &str { + "stability" + } + + fn provider_name(&self) -> &str { + "Stability AI" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 图像生成:`POST /v1/generation/{engine_id}/text-to-image` + /// + /// `engine_id` 取 `req.model`,为空时用 [`DEFAULT_ENGINE`] 兜底。 + /// 响应解析 base64 图像到 `ImageResult.data`。 + async fn image_generate(&self, req: ImageRequest) -> Result { + self.ensure_capability(Capabilities::ImageGenerate)?; + let engine = if req.model.is_empty() { + DEFAULT_ENGINE.to_string() + } else { + req.model.clone() + }; + let body = self.build_generate_body(&req); + let path = format!("v1/generation/{engine}/text-to-image"); + let value = self.post_authed_json(&path, &body).await?; + Self::parse_image_result(&value, &engine) + } + + /// 模型列表(实时拉取 Stability engines) + /// + /// `GET /v1/engines/list`,返回顶层数组,按 `filter` 过滤模型类型。 + async fn list_models(&self, filter: Option) -> Result> { + let value = self.get_authed_json("v1/engines/list").await?; + let models = Self::parse_engines(&value); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / video / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 辅助函数 ==================== + +/// 从 ImageRequest 的 width/height/size 解析图像尺寸 +/// +/// 优先级:显式 width+height > size 字符串 > 默认 1024x1024。 +/// 仅提供 width 或 height 之一时,另一维度用默认值。 +/// size 字符串解析失败时回退到默认值。 +/// +/// 与 Python v1 `kwargs.get("width", 1024)` / `kwargs.get("height", 1024)` 行为一致。 +fn resolve_dimensions(width: Option, height: Option, size: Option<&str>) -> (u32, u32) { + match (width, height) { + (Some(w), Some(h)) => (w, h), + (Some(w), None) => (w, DEFAULT_HEIGHT), + (None, Some(h)) => (DEFAULT_WIDTH, h), + (None, None) => match size.and_then(|s| util::parse_size(s).ok()) { + Some((w, h)) => (w, h), + None => (DEFAULT_WIDTH, DEFAULT_HEIGHT), + }, + } +} + +/// 解析单个 artifact(Stability 图像响应项) +/// +/// Stability artifact 字段:`base64`(图像数据)/ `seed` / `finishReason`。 +/// 仅提取 base64,映射到 `ImageData.b64_json`。base64 为空的项被过滤。 +fn parse_artifact(v: &Value) -> Option { + let b64 = v.get("base64").and_then(|x| x.as_str())?; + if b64.is_empty() { + return None; + } + Some(ImageData { + b64_json: Some(b64.to_owned()), + url: None, + revised_prompt: None, + }) +} + +/// 将 reqwest::Error 映射为 AibridgeError +/// +/// 超时 → Timeout;其余 → Network。 +fn map_reqwest_error(err: reqwest::Error) -> AibridgeError { + if err.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(err) + } +} + +/// 解析错误体中的 message 字段 +/// +/// 兼容多种 Stability 错误体结构: +/// - `{"error": {"message": "..."}}`(OpenAI 风格) +/// - `{"message": "..."}`(Stability 常见顶层 message) +/// - `{"errors": [{"message": "..."}]}`(Stability 校验错误数组) +/// - `{"error": "..."}`(顶层 error 字符串) +/// +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + // errors 数组(Stability validation 错误) + if let Some(arr) = v.get("errors").and_then(|m| m.as_array()) { + let parts: Vec = arr + .iter() + .filter_map(|x| x.get("message").and_then(|m| m.as_str()).map(str::to_owned)) + .collect(); + if !parts.is_empty() { + return parts.join("; "); + } + } + // error.message(嵌套) + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 顶层 error 字符串 + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::video::VideoRequest; + use mockito::Server; + use serde_json::json; + + // ==================== 测试辅助 ==================== + + /// 构造测试用 StabilityAdapter(指向 mockito server) + fn make_adapter(server: &Server) -> StabilityAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("stability", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + StabilityAdapter::with_http(http, config) + } + + /// 构造不指向任何 server 的 StabilityAdapter(用于不发请求的元信息/能力测试) + fn make_adapter_no_server() -> StabilityAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_STABILITY_BASE_URL) + .build(); + let config = ProviderConfig::from_options("stability", opts); + StabilityAdapter::new(config).expect("StabilityAdapter 构造应成功") + } + + // ============ 元信息 ============ + + #[test] + fn provider_type_and_name_match_python() { + let adapter = make_adapter_no_server(); + assert_eq!(adapter.provider_type(), "stability"); + assert_eq!(adapter.provider_name(), "Stability AI"); + } + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_only_image() { + let adapter = make_adapter_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::ImageGenerate)); + // chat / video 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::VideoGenerate)); + } + + #[test] + fn base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("stability", opts); + let adapter = StabilityAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_STABILITY_BASE_URL); + } + + #[test] + fn base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.stability-proxy.com") + .build(); + let config = ProviderConfig::from_options("stability", opts); + let adapter = StabilityAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.stability-proxy.com"); + } + + // ============ image_generate 正常路径 ============ + + #[tokio::test] + async fn image_generate_success_parses_base64() { + let mut server = Server::new_async().await; + let body = json!({ + "artifacts": [ + {"base64": "aGVsbG8=", "finishReason": "SUCCESS", "seed": 123}, + {"base64": "d29ybGQ=", "finishReason": "SUCCESS", "seed": 456} + ] + }); + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_header("authorization", "Bearer test-key") + .match_header("accept", "application/json") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let resp = adapter + .image_generate(req) + .await + .expect("image_generate 应成功"); + + assert_eq!(resp.model, "stable-diffusion-xl-1024-v1-1"); + assert_eq!(resp.object, "image.generation"); + assert_eq!(resp.data.len(), 2); + assert_eq!(resp.data[0].b64_json.as_deref(), Some("aGVsbG8=")); + assert_eq!(resp.data[1].b64_json.as_deref(), Some("d29ybGQ=")); + // Stability 只返回 base64,无 url + assert!(resp.data[0].url.is_none()); + // id 与 created 应已填充 + assert!(resp.id.starts_with("img_")); + assert!(resp.created > 0); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_uses_default_engine_when_model_empty() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + // 回退到默认引擎 + assert_eq!(resp.model, DEFAULT_ENGINE); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_uses_custom_engine_in_path() { + // 验证自定义 model 作为 engine_id 出现在 URL 路径 + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-3-medium/text-to-image", + ) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-3-medium", "a cat").build(); + let resp = adapter.image_generate(req).await.unwrap(); + assert_eq!(resp.model, "stable-diffusion-3-medium"); + mock.assert_async().await; + } + + // ============ image_generate 参数映射 ============ + + #[tokio::test] + async fn image_generate_sends_text_prompts_with_weight() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "text_prompts": [ + {"text": "a cat", "weight": 1}, + {"text": "blurry", "weight": -1} + ] + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .negative_prompt("blurry") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_default_size_1024x1024() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "width": 1024, + "height": 1024 + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_parses_size_string_to_width_height() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "width": 1344, + "height": 768 + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .size("1344x768") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_width_height_take_priority_over_size() { + // 同时提供 width/height 和 size,应优先用 width/height + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "width": 1536, + "height": 640 + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .width(1536) + .height(640) + .size("1024x1024") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_passes_sampling_params() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "steps": 40, + "seed": 42, + "cfg_scale": 7.5 + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .steps(40) + .seed(42) + .cfg_scale(7.5) + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_caps_samples_at_10() { + // 请求 n=20,应被截断为 samples=10 + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "samples": 10 + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .n(20) + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_passes_style_preset() { + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "style_preset": "anime" + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .style("anime") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_passes_extra_params() { + // extra 字段透传到顶层(如 sampler、sampler_seed 等厂商特有参数) + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "sampler": "K_DPMPP_2M", + "custom_param": "value" + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat") + .extra("sampler", "K_DPMPP_2M") + .extra("custom_param", "value") + .build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn image_generate_no_negative_prompt_omits_second_entry() { + // 无 negative_prompt 时 text_prompts 仅含正向一项 + let mut server = Server::new_async().await; + let mock = server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .match_body(mockito::Matcher::PartialJson(json!({ + "text_prompts": [{"text": "a cat", "weight": 1}] + }))) + .with_status(200) + .with_body(json!({"artifacts": [{"base64": "aGk="}]}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let _ = adapter.image_generate(req).await.unwrap(); + mock.assert_async().await; + } + + // ============ image_generate 错误路径 ============ + + #[tokio::test] + async fn image_generate_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(401) + .with_body(json!({"message": "invalid api key"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn image_generate_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(429) + .with_body(json!({"message": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn image_generate_error_400_returns_validation() { + // 任务要求 400 → Validation(携带 details) + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(400) + .with_body(json!({"message": "invalid width"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, details } => { + assert_eq!(message, "invalid width"); + // details 携带原始错误体 + assert_eq!(details["message"], "invalid width"); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn image_generate_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(500) + .with_body(json!({"message": "internal error"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal error"); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn image_generate_error_403_returns_api() { + // 403(引擎不可用/配额)走 Api(与 Python v1 一致) + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(403) + .with_body(json!({"message": "engine not accessible"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 403); + assert!(message.contains("engine not accessible")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn image_generate_error_no_json_body_falls_back() { + // 非 JSON 错误体回退到 HTTP {status} + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(502) + .with_body("Bad Gateway") + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn image_generate_no_artifacts_returns_api_error() { + // artifacts 为空时返 Api 错误(与 Python v1 `raise APIError("No image generated")` 一致) + let mut server = Server::new_async().await; + server + .mock( + "POST", + "/v1/generation/stable-diffusion-xl-1024-v1-1/text-to-image", + ) + .with_status(200) + .with_body(json!({"artifacts": []}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let req = ImageRequest::builder("stable-diffusion-xl-1024-v1-1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("artifacts")), + _ => panic!("应为 Api"), + } + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_success_parses_top_level_array() { + // Stability engines/list 返回顶层数组 + let mut server = Server::new_async().await; + let mock = server + .mock("GET", "/v1/engines/list") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body( + json!([ + {"id": "stable-diffusion-xl-1024-v1-1", "name": "Stable Diffusion XL 1.1", "description": "SDXL", "create_time": 1700000000}, + {"id": "stable-diffusion-3-medium", "name": "Stable Diffusion 3 Medium", "create_time": 1700000100} + ]) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "stable-diffusion-xl-1024-v1-1"); + assert_eq!(models[0].name, "Stable Diffusion XL 1.1"); + assert_eq!(models[0].provider, "stability"); + // stable-diffusion 关键字 → Image + assert_eq!(models[0].model_type, ModelType::Image); + assert_eq!(models[0].description.as_deref(), Some("SDXL")); + assert_eq!(models[0].created, Some(1700000000)); + assert!(!models[0].supports_streaming); + mock.assert_async().await; + } + + #[tokio::test] + async fn list_models_filter_by_image_type() { + let mut server = Server::new_async().await; + server + .mock("GET", "/v1/engines/list") + .with_status(200) + .with_body( + json!([ + {"id": "stable-diffusion-xl-1024-v1-1", "name": "SDXL"}, + {"id": "esrgan-v1-x2plus", "name": "ESRGAN"} + ]) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + // 仅 stable-diffusion 命中 Image 关键字 + assert_eq!(images.len(), 1); + assert_eq!(images[0].id, "stable-diffusion-xl-1024-v1-1"); + } + + #[tokio::test] + async fn list_models_accepts_data_wrapper() { + // 兼容 {"data": [...]} 包装格式 + let mut server = Server::new_async().await; + server + .mock("GET", "/v1/engines/list") + .with_status(200) + .with_body( + json!({ + "data": [{"id": "stable-diffusion-xl-1024-v1-1", "name": "SDXL"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "stable-diffusion-xl-1024-v1-1"); + } + + #[tokio::test] + async fn list_models_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/v1/engines/list") + .with_status(401) + .with_body(json!({"message": "bad key"}).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let err = adapter.list_models(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_models_empty_array() { + let mut server = Server::new_async().await; + server + .mock("GET", "/v1/engines/list") + .with_status(200) + .with_body(json!([]).to_string()) + .create_async() + .await; + + let adapter = make_adapter(&server); + let models = adapter.list_models(None).await.unwrap(); + assert!(models.is_empty()); + } + + // ============ 不支持的能力 ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter_no_server(); + let req = ChatRequest::builder("m", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter_no_server(); + let req = VideoRequest::builder("m", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn video_poll_returns_unsupported() { + let adapter = make_adapter_no_server(); + let err = adapter.video_poll("task-1", "m").await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ start / close ============ + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ map_api_error 单元测试 ============ + + #[test] + fn map_api_error_401_is_authentication() { + let err = StabilityAdapter::map_api_error(401, ""); + match err { + AibridgeError::Authentication { message } => assert!(message.contains("Stability")), + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn map_api_error_429_is_rate_limit() { + let err = StabilityAdapter::map_api_error(429, ""); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn map_api_error_400_is_validation_with_details() { + let body = json!({"message": "invalid width"}).to_string(); + let err = StabilityAdapter::map_api_error(400, &body); + match err { + AibridgeError::Validation { message, details } => { + assert_eq!(message, "invalid width"); + assert_eq!(details["message"], "invalid width"); + } + _ => panic!("应为 Validation"), + } + } + + #[test] + fn map_api_error_400_extracts_errors_array() { + // Stability 校验错误数组格式 + let body = + json!({"errors": [{"message": "field1 invalid"}, {"message": "field2 invalid"}]}) + .to_string(); + let err = StabilityAdapter::map_api_error(400, &body); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("field1 invalid")); + assert!(message.contains("field2 invalid")); + } + _ => panic!("应为 Validation"), + } + } + + #[test] + fn map_api_error_500_is_api() { + let err = StabilityAdapter::map_api_error(500, "{\"message\":\"internal\"}"); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_403_is_api() { + let err = StabilityAdapter::map_api_error(403, "{\"message\":\"forbidden\"}"); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 403); + assert_eq!(message, "forbidden"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_falls_back_to_http_status() { + let err = StabilityAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ resolve_dimensions 单元测试 ============ + + #[test] + fn resolve_dimensions_uses_width_height_when_both_present() { + assert_eq!(resolve_dimensions(Some(512), Some(512), None), (512, 512)); + assert_eq!(resolve_dimensions(Some(1344), Some(768), None), (1344, 768)); + } + + #[test] + fn resolve_dimensions_uses_size_string_when_no_width_height() { + assert_eq!( + resolve_dimensions(None, None, Some("1024x1024")), + (1024, 1024) + ); + assert_eq!( + resolve_dimensions(None, None, Some("1344x768")), + (1344, 768) + ); + } + + #[test] + fn resolve_dimensions_defaults_when_nothing_provided() { + assert_eq!(resolve_dimensions(None, None, None), (1024, 1024)); + } + + #[test] + fn resolve_dimensions_falls_back_when_size_invalid() { + // size 字符串非法时回退到默认 1024x1024 + assert_eq!(resolve_dimensions(None, None, Some("abc")), (1024, 1024)); + assert_eq!(resolve_dimensions(None, None, Some("1024")), (1024, 1024)); + } + + #[test] + fn resolve_dimensions_partial_width_uses_default_height() { + assert_eq!(resolve_dimensions(Some(512), None, None), (512, 1024)); + } + + #[test] + fn resolve_dimensions_partial_height_uses_default_width() { + assert_eq!(resolve_dimensions(None, Some(512), None), (1024, 512)); + } + + // ============ parse_engines / parse_image_result 单元测试 ============ + + #[test] + fn parse_engines_top_level_array() { + let value = json!([ + {"id": "stable-diffusion-xl-1024-v1-1", "name": "SDXL"}, + {"id": "stable-diffusion-3-medium", "name": "SD3"} + ]); + let models = StabilityAdapter::parse_engines(&value); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "stable-diffusion-xl-1024-v1-1"); + assert_eq!(models[0].model_type, ModelType::Image); + } + + #[test] + fn parse_engines_data_wrapper() { + let value = json!({"data": [{"id": "stable-diffusion-xl-1024-v1-1"}]}); + let models = StabilityAdapter::parse_engines(&value); + assert_eq!(models.len(), 1); + } + + #[test] + fn parse_engines_empty_array() { + let value = json!([]); + let models = StabilityAdapter::parse_engines(&value); + assert!(models.is_empty()); + } + + #[test] + fn parse_engines_uses_id_as_name_when_missing() { + let value = json!([{"id": "stable-diffusion-xl-1024-v1-1"}]); + let models = StabilityAdapter::parse_engines(&value); + assert_eq!(models[0].name, "stable-diffusion-xl-1024-v1-1"); + } + + #[test] + fn parse_image_result_extracts_base64() { + let value = json!({ + "artifacts": [{"base64": "aGk="}, {"base64": "eQ=="}] + }); + let result = StabilityAdapter::parse_image_result(&value, "sdxl").unwrap(); + assert_eq!(result.data.len(), 2); + assert_eq!(result.data[0].b64_json.as_deref(), Some("aGk=")); + assert_eq!(result.model, "sdxl"); + } + + #[test] + fn parse_image_result_empty_artifacts_returns_error() { + let value = json!({"artifacts": []}); + let result = StabilityAdapter::parse_image_result(&value, "sdxl"); + assert!(result.is_err()); + } + + #[test] + fn parse_image_result_missing_artifacts_returns_error() { + let value = json!({}); + let result = StabilityAdapter::parse_image_result(&value, "sdxl"); + assert!(result.is_err()); + } + + #[test] + fn parse_image_result_filters_empty_base64() { + // base64 为空的 artifact 被过滤;全部为空则返错 + let value = json!({"artifacts": [{"base64": ""}, {"base64": "aGk="}]}); + let result = StabilityAdapter::parse_image_result(&value, "sdxl").unwrap(); + assert_eq!(result.data.len(), 1); + assert_eq!(result.data[0].b64_json.as_deref(), Some("aGk=")); + } + + // ============ build_generate_body 单元测试 ============ + + #[test] + fn build_generate_body_includes_required_fields() { + let adapter = make_adapter_no_server(); + let req = ImageRequest::builder("sdxl", "a cat") + .negative_prompt("blurry") + .build(); + let body = adapter.build_generate_body(&req); + // text_prompts 含正向 + 负向 + let prompts = body.get("text_prompts").and_then(|v| v.as_array()).unwrap(); + assert_eq!(prompts.len(), 2); + assert_eq!(prompts[0]["text"], "a cat"); + assert_eq!(prompts[0]["weight"], 1); + assert_eq!(prompts[1]["text"], "blurry"); + assert_eq!(prompts[1]["weight"], -1); + // 默认尺寸 + assert_eq!(body["width"], 1024); + assert_eq!(body["height"], 1024); + // samples 默认 1 + assert_eq!(body["samples"], 1); + } + + #[test] + fn build_generate_body_merges_extra() { + let adapter = make_adapter_no_server(); + let req = ImageRequest::builder("sdxl", "a cat") + .extra("sampler", "K_DPMPP_2M") + .build(); + let body = adapter.build_generate_body(&req); + assert_eq!(body["sampler"], "K_DPMPP_2M"); + } +} From 04fd2027a2ebdaec4e6a4d7cbb7d74f3bfe5c335 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 00:47:49 +0800 Subject: [PATCH 35/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20=E7=AC=AC=E4=B8=80=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20anthropic=20+=20stability=20=E5=88=B0=E5=B7=A5?= =?UTF-8?q?=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 36 ++++++++++++++++++--- crates/aibridge-core/src/adapters/mod.rs | 6 ++++ 2 files changed, 37 insertions(+), 5 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index eaf9a00..e905ff6 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -15,6 +15,7 @@ use crate::adapters::aggregation_platforms::{ CloudflareAIAdapter, FireworksAIAdapter, SiliconFlowAdapter, TogetherAIAdapter, }; use crate::adapters::agnes::AgnesAdapter; +use crate::adapters::anthropic::AnthropicAdapter; use crate::adapters::azure::AzureAdapter; use crate::adapters::chinese::{ DoubaoAdapter, ErnieAdapter, KimiAdapter, MiniMaxAdapter, QwenAdapter, ZhipuAdapter, @@ -26,6 +27,7 @@ use crate::adapters::more_models::{ CohereAdapter, DeepSeekAdapter, MistralAdapter, PerplexityAdapter, StepFunAdapter, }; use crate::adapters::openai::OpenAiAdapter; +use crate::adapters::stability::StabilityAdapter; use crate::adapters::volcengine_cv::VolcengineCvAdapter; use crate::config::ProviderConfig; use crate::error::{AibridgeError, Result}; @@ -67,12 +69,13 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "ernie", "kimi", "minimax", - // 阶段 2b/2c 待实现: + // 阶段 2b 第一批独立协议(已实现): "anthropic", + "stability", + // 阶段 2b/2c 待实现: "runway", "pika", "kling", - "stability", "edge-tts", "elevenlabs", "cartesia", @@ -128,9 +131,12 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { "ernie" => Ok(Box::new(ErnieAdapter::new(config)?)), "kimi" => Ok(Box::new(KimiAdapter::new(config)?)), "minimax" => Ok(Box::new(MiniMaxAdapter::new(config)?)), + // 阶段 2b 独立协议:别名对齐 Python agn/adapters/{anthropic,stability}.py 末尾 register 调用(均无别名) + "anthropic" => Ok(Box::new(AnthropicAdapter::new(config)?)), + "stability" => Ok(Box::new(StabilityAdapter::new(config)?)), // 阶段 2 适配器占位 - "anthropic" | "runway" | "pika" | "kling" | "stability" | "edge-tts" | "elevenlabs" - | "cartesia" | "deepgram" | "assemblyai" => Err(AibridgeError::ProviderNotFound { + "runway" | "pika" | "kling" | "edge-tts" | "elevenlabs" | "cartesia" | "deepgram" + | "assemblyai" => Err(AibridgeError::ProviderNotFound { provider: format!("{provider}(阶段 2 待实现)"), }), // 未知 provider @@ -400,6 +406,22 @@ mod tests { assert_eq!(adapter.provider_type(), "minimax"); } + #[test] + fn create_anthropic_returns_adapter() { + // 阶段 2b:AnthropicAdapter 自带 DEFAULT_ANTHROPIC_BASE_URL 兜底,仅需 api_key + let adapter = + create_adapter(config_for("anthropic")).expect("工厂应能创建 anthropic 适配器"); + assert_eq!(adapter.provider_type(), "anthropic"); + } + + #[test] + fn create_stability_returns_adapter() { + // 阶段 2b:StabilityAdapter 自带 DEFAULT_STABILITY_BASE_URL 兜底,仅需 api_key + let adapter = + create_adapter(config_for("stability")).expect("工厂应能创建 stability 适配器"); + assert_eq!(adapter.provider_type(), "stability"); + } + #[test] fn create_additional_models_aliases_map_to_main_provider_type() { // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: @@ -427,7 +449,8 @@ mod tests { #[test] fn create_phase2_adapter_returns_phase2_message() { - let result = create_adapter(config_for("anthropic")); + // runway 仍为阶段 2 占位(未实现),返 ProviderNotFound + let result = create_adapter(config_for("runway")); if let Err(AibridgeError::ProviderNotFound { provider }) = result { assert!(provider.contains("阶段 2")); } else { @@ -466,6 +489,9 @@ mod tests { assert!(is_known_provider("ernie")); assert!(is_known_provider("kimi")); assert!(is_known_provider("minimax")); + // 阶段 2b 第一批独立协议 + assert!(is_known_provider("anthropic")); + assert!(is_known_provider("stability")); } #[test] diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index c3f25f1..5062c1b 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -45,3 +45,9 @@ pub mod emerging_models; /// 中文模型适配器:阶段 2a,含 Qwen/Zhipu/Doubao/Ernie/Kimi/MiniMax 六个 OpenAI 兼容子适配器 pub mod chinese; + +/// Anthropic Claude 适配器:阶段 2b 独立协议,文本对话/流式/多模态 +pub mod anthropic; + +/// Stability AI 适配器:阶段 2b 独立协议,文生图/图生图 +pub mod stability; From 0247589dbf288e28abd24cec13fc06004cbb080d Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 00:58:53 +0800 Subject: [PATCH 36/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20runway=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/runway.rs | 1553 +++++++++++++++++++ 1 file changed, 1553 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/runway.rs diff --git a/crates/aibridge-core/src/adapters/runway.rs b/crates/aibridge-core/src/adapters/runway.rs new file mode 100644 index 0000000..84a2d6e --- /dev/null +++ b/crates/aibridge-core/src/adapters/runway.rs @@ -0,0 +1,1553 @@ +//! Runway 视频生成适配器 +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/runway.py`。 +//! +//! Runway 为**独立协议**(非 OpenAI 兼容),不复用 `OpenAiCompatAdapter`: +//! - 文生视频:`POST /text_to_video` +//! - 图生视频:`POST /image_to_video`(需 `promptImage`) +//! - 查询任务:`GET /assets/{id}` +//! - 认证:`Authorization: Bearer {api_key}` +//! - 默认 Base URL:`https://api.runwayml.com/v1` +//! +//! ## 阶段范围 +//! 阶段 2b 实现 video_create + video_poll + list_models(硬编码)。 +//! 不支持的能力(chat / image / embed / transcribe / speech)走 `Adapter` trait +//! 默认实现返 `UnsupportedCapability`,与 Python v1 抛 `UnsupportedCapabilityError` 行为一致。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::common::{ModelInfo, ModelType, TaskStatus, VideoMode}; +use crate::model::image::FileInput; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +/// Runway 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +pub const DEFAULT_RUNWAY_BASE_URL: &str = "https://api.runwayml.com/v1"; + +/// Runway 支持的能力集合 +/// +/// 对齐 Python v1 `RunwayAdapter.supported_capabilities = ["video"]`。 +/// Rust 用 `VideoGenerate` 表达视频生成能力(含 video_create + video_poll), +/// 并声明 text2video / image2video 两个子能力。 +fn runway_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps +} + +/// Runway 适配器 +/// +/// Runway Gen-3 Alpha / Gen-3 Turbo 视频生成平台,支持文生视频与图生视频。 +/// 官方 API 文档:https://docs.dev.runwayml.ai/ +/// +/// ## API 规范 +/// - Base URL: `https://api.runwayml.com/v1` +/// - 文生视频: `POST /text_to_video`(body 含 `prompt`) +/// - 图生视频: `POST /image_to_video`(body 含 `promptImage` + `promptText`) +/// - 查询任务: `GET /assets/{id}` +/// - 认证: Bearer Token +pub struct RunwayAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl RunwayAdapter { + /// 创建 Runway 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_RUNWAY_BASE_URL`] 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_RUNWAY_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: runway_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_RUNWAY_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: runway_capabilities(), + } + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: runway)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用 Runway 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 Bearer 认证的 GET 请求,并用 Runway 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 Runway 视频创建请求体,并返回对应端点 + /// + /// 移植自 Python v1 `video_create`: + /// - `image2video` 模式且有参考图 → `POST /image_to_video`,body 含 `promptImage` + `promptText` + /// - 其余 → `POST /text_to_video`,body 含 `prompt` + /// - 参考图优先 `reference_images[0]`,回退 `first_frame`(Runway 仅支持单张首帧) + /// - 可选参数:width / height / seed / motion / cameraMotion / aspectRatio + /// - extra 字段合并到顶层(透传厂商特有参数) + fn build_create_body(&self, req: &VideoRequest) -> (&'static str, Value) { + // 默认模型 gen-3(与 list_models 一致) + let model = if req.model.is_empty() { + "gen-3".to_string() + } else { + req.model.clone() + }; + + // 图生视频:优先 reference_images[0],回退 first_frame + let prompt_image = req + .reference_images + .first() + .map(file_input_to_url) + .filter(|s| !s.is_empty()) + .or_else(|| req.first_frame.as_ref().map(file_input_to_url)) + .filter(|s| !s.is_empty()); + + let is_image2video = matches!(req.mode, VideoMode::Image2Video); + + let (endpoint, mut body) = match (is_image2video, prompt_image) { + // 图生视频:Runway 用 promptImage + promptText + (true, Some(img)) => { + let b = json!({ + "model": model, + "promptImage": img, + "promptText": req.prompt, + }); + ("image_to_video", b) + } + // 文生视频:Runway 用 prompt + _ => { + let b = json!({ + "model": model, + "prompt": req.prompt, + }); + ("text_to_video", b) + } + }; + + // 可选参数(移植自 Python v1:kwargs 有则加) + body["width"] = json!(req.width); + body["height"] = json!(req.height); + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + if let Some(motion) = req.motion_strength { + body["motion"] = json!(motion); + } + if let Some(cm) = &req.camera_motion { + body["cameraMotion"] = json!(cm); + } + if let Some(ar) = &req.aspect_ratio { + body["aspectRatio"] = json!(ar); + } + + // extra 透传(合并到顶层,覆盖同名字段) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + + (endpoint, body) + } + + /// 解析 Runway 创建任务响应 → VideoTask + /// + /// Runway 创建响应:`{"id": "...", "status": "pending", "createdAt": "..."}` + /// task_id 兼容 `id` / `assetId` / `taskId` 字段,缺失时生成 `vid` 前缀 ID(与 Python 一致)。 + fn parse_video_task(value: &Value, model: &str) -> Result { + let task_id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("assetId") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("taskId") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .unwrap_or_else(|| util::generate_id("vid")); + let raw_status = value + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("pending"); + Ok(VideoTask { + task_id, + model: model.to_string(), + status: map_runway_status(raw_status), + created_at: util::current_timestamp(), + }) + } + + /// 解析 Runway 任务查询响应 → VideoStatus + /// + /// 移植自 Python v1 `video_poll`: + /// - 视频 URL 兼容 `url` / `videoUrl` / `output.url` / `output.videoUrl` / `assets[0].url` + /// - 错误信息兼容 `error` / `errorMessage` + /// - progress / createdAt / updatedAt 直接透传(数字形式) + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let raw_status = value.get("status").and_then(|v| v.as_str()).unwrap_or(""); + let status = map_runway_status(raw_status); + + // 视频 URL:多路径兼容(与 Python v1 的提取顺序一致) + let video_url = value + .get("url") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("videoUrl") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("videoUrl")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("assets") + .and_then(|a| a.as_array()) + .and_then(|arr| arr.first()) + .and_then(|item| item.get("url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + + // 错误信息:error / errorMessage + let error = value + .get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("errorMessage") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + + let progress = value + .get("progress") + .and_then(|v| v.as_u64()) + .map(|p| p as u32); + + // createdAt / updatedAt:Runway 通常返回 ISO 字符串,此处仅在为数字时解析 + // (与方舟 volcengine_cv 行为一致;ISO 字符串场景留待后续统一重构) + let created_at = value.get("createdAt").and_then(|v| v.as_u64()); + let updated_at = value.get("updatedAt").and_then(|v| v.as_u64()); + + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress, + error, + created_at, + updated_at, + } + } + + /// 将 Runway API 错误响应映射为 AibridgeError + /// + /// 映射规则(与 [`AibridgeError::from_http_status`] 一致): + /// - 401/403 → Authentication("Invalid Runway API key") + /// - 429 → RateLimit("Runway rate limit exceeded") + /// - 404 → ModelNotFound(提取错误体消息) + /// - 400 → Validation(提取错误体消息,details 携带原始 body) + /// - 其余 ≥400(含 5xx)→ Api + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 | 403 => AibridgeError::Authentication { + message: "Invalid Runway API key".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Runway rate limit exceeded".to_string(), + retry_after: None, + }, + 404 => AibridgeError::ModelNotFound { + model: parse_error_message(body, status), + }, + 400 => AibridgeError::Validation { + message: parse_error_message(body, status), + details: serde_json::from_str::(body).unwrap_or(Value::Null), + }, + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for RunwayAdapter { + fn provider_type(&self) -> &str { + "runway" + } + + fn provider_name(&self) -> &str { + "Runway" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 创建视频生成任务 + /// + /// 根据模式与参考图选择端点: + /// - `image2video` 且有参考图 → `POST /image_to_video` + /// - 其余 → `POST /text_to_video` + async fn video_create(&self, req: VideoRequest) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let model = if req.model.is_empty() { + "gen-3".to_string() + } else { + req.model.clone() + }; + let (endpoint, body) = self.build_create_body(&req); + let value = self.post_authed_json(endpoint, &body).await?; + Self::parse_video_task(&value, &model) + } + + /// 查询视频任务状态:`GET /assets/{task_id}` + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let path = format!("assets/{task_id}"); + let value = self.get_authed_json(&path).await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 模型列表(硬编码) + /// + /// Runway 无标准 `/models` 端点,暂保留硬编码列表(与 Python v1 一致)。 + /// 含 gen-3 / gen-3-turbo 两个视频模型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = runway_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / image_generate / embed / transcribe / speech 走 trait 默认实现 + // 返 UnsupportedCapability,与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 内部:硬编码模型列表 ==================== + +/// Runway 硬编码模型列表 +/// +/// 对应 Python v1 `RunwayAdapter.list_models`。 +/// 注意:该 Provider 无标准 `/models` 端点,暂保留硬编码列表。 +fn runway_hardcoded_models() -> Vec { + vec![ + ModelInfo { + id: "gen-3".into(), + name: "Gen-3 Alpha".into(), + model_type: ModelType::Video, + provider: "runway".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Runway Gen-3 Alpha 视频生成模型".into()), + created: None, + }, + ModelInfo { + id: "gen-3-turbo".into(), + name: "Gen-3 Turbo".into(), + model_type: ModelType::Video, + provider: "runway".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Runway Gen-3 Turbo 快速视频生成模型".into()), + created: None, + }, + ] +} + +// ==================== 辅助函数 ==================== + +/// 映射 Runway 状态字符串到统一 TaskStatus +/// +/// 移植自 Python v1 `_map_runway_status`: +/// - pending / queued → Pending +/// - processing / running / in_progress → Processing +/// - completed / succeeded / success → Success +/// - failed / error / cancelled → Failed +/// - 未知 → Pending(与 Python 默认值一致) +fn map_runway_status(raw: &str) -> TaskStatus { + match raw.to_lowercase().as_str() { + "pending" | "queued" => TaskStatus::Pending, + "processing" | "running" | "in_progress" => TaskStatus::Processing, + "completed" | "succeeded" | "success" => TaskStatus::Success, + "failed" | "error" | "cancelled" => TaskStatus::Failed, + _ => TaskStatus::Pending, + } +} + +/// 从 FileInput 提取 URL/base64 字符串 +/// +/// Runway `promptImage` 需要 URL 或 base64 字符串: +/// - `Url(s)` / `Base64(s)` → `s` +/// - `Path(_)` / `Bytes(_)` → 空字符串(需上层预转换为 URL/base64,此处不处理) +fn file_input_to_url(input: &FileInput) -> String { + match input { + FileInput::Url(s) | FileInput::Base64(s) => s.clone(), + FileInput::Path(_) | FileInput::Bytes(_) => String::new(), + } +} + +/// 解析错误体中的 message 字段 +/// +/// 兼容多种错误体结构(Runway 无统一规范,多种均可能出现): +/// - `{"error": "..."}`(顶层 error 字符串) +/// - `{"error": {"message": "..."}}`(OpenAI 风格) +/// - `{"message": "..."}` +/// - `{"detail": "..."}` +/// +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + // 顶层 error 字符串(Runway 常见) + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // OpenAI 风格 error.message + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 detail + if let Some(msg) = v.get("detail").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::http::HttpClient; + use crate::model::chat::{ChatMessage, ChatRequest}; + use crate::model::image::ImageRequest; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 RunwayAdapter(指向 mockito server) + fn make_runway(server: &Server) -> RunwayAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("runway", opts); + // http 客户端的 base_url 指向 mockito(与 emerging_models 测试结构一致) + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + RunwayAdapter::with_http(http, config) + } + + /// 构造不指向任何 server 的 RunwayAdapter(用于不发请求的元信息/能力测试) + fn make_runway_no_server() -> RunwayAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_RUNWAY_BASE_URL) + .build(); + let config = ProviderConfig::from_options("runway", opts); + RunwayAdapter::new(config).expect("RunwayAdapter 构造应成功") + } + + // ============ 元信息 ============ + + #[test] + fn runway_provider_type_and_name_match_python() { + let adapter = make_runway_no_server(); + assert_eq!(adapter.provider_type(), "runway"); + assert_eq!(adapter.provider_name(), "Runway"); + } + + #[test] + fn runway_requires_api_key_is_true() { + let adapter = make_runway_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn runway_capabilities_contains_only_video() { + let adapter = make_runway_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + // chat / image 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn runway_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("runway", opts); + let adapter = RunwayAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_RUNWAY_BASE_URL); + } + + #[test] + fn runway_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.runway-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("runway", opts); + let adapter = RunwayAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.runway-proxy.com/v1"); + } + + // ============ video_create 正常路径 ============ + + #[tokio::test] + async fn runway_video_create_text2video_success_returns_task() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "asset-abc", + "status": "pending", + "createdAt": "2024-01-01T00:00:00Z" + }); + let mock = server + .mock("POST", "/text_to_video") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat walking") + .aspect_ratio("16:9") + .seed(42) + .build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + + assert_eq!(task.task_id, "asset-abc"); + assert_eq!(task.model, "gen-3"); + assert_eq!(task.status, TaskStatus::Pending); + assert!(task.created_at > 0); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_text2video_sends_model_and_prompt() { + // 验证 text_to_video 端点的请求体结构 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/text_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gen-3", + "prompt": "a cat", + "width": 1280, + "height": 720, + "seed": 42, + "aspectRatio": "16:9" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat") + .aspect_ratio("16:9") + .seed(42) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_uses_default_model_when_empty() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/text_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gen-3" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.model, "gen-3"); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_image2video_with_reference_image() { + // image2video 模式:走 /image_to_video,body 含 promptImage + promptText + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/image_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "gen-3", + "promptImage": "https://example.com/start.png", + "promptText": "animate this" + }))) + .with_status(200) + .with_body(json!({"id": "img-task", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "animate this") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/start.png")]) + .build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "img-task"); + assert_eq!(task.status, TaskStatus::Pending); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_image2video_falls_back_to_first_frame() { + // image2video 但无 reference_images,用 first_frame 作为 promptImage + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/image_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "promptImage": "https://example.com/first.png" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "animate") + .mode(VideoMode::Image2Video) + .first_frame(FileInput::url("https://example.com/first.png")) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_image2video_without_image_falls_back_to_text() { + // image2video 但既无 reference_images 也无 first_frame → 回退 text_to_video + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/text_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "prompt": "a cat" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat") + .mode(VideoMode::Image2Video) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_passes_motion_and_camera_motion() { + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/text_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "motion": 7.5, + "cameraMotion": "zoom_in" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat") + .motion_strength(7.5) + .camera_motion("zoom_in") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_passes_extra_params() { + // extra 透传到顶层 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/text_to_video") + .match_body(mockito::Matcher::PartialJson(json!({ + "custom_param": "custom_value", + "watermark": false + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat") + .extra("custom_param", "custom_value") + .extra("watermark", false) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_create_uses_assetid_field_when_id_missing() { + // 响应缺 id 但有 assetId + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(200) + .with_body(json!({"assetId": "asset-xyz", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "asset-xyz"); + } + + #[tokio::test] + async fn runway_video_create_generates_id_when_all_id_fields_missing() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(200) + .with_body(json!({"status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert!(task.task_id.starts_with("vid_")); + } + + // ============ video_create 错误路径 ============ + + #[tokio::test] + async fn runway_video_create_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(401) + .with_body(json!({"error": "Invalid API key"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn runway_video_create_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(429) + .with_body(json!({"error": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn runway_video_create_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(404) + .with_body(json!({"error": "model not found"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-unknown", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn runway_video_create_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(400) + .with_body(json!({"error": "invalid prompt"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("invalid prompt")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn runway_video_create_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(500) + .with_body(json!({"error": "internal"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert!(message.contains("internal")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn runway_video_create_error_no_json_body_falls_back() { + let mut server = Server::new_async().await; + server + .mock("POST", "/text_to_video") + .with_status(502) + .with_body("Bad Gateway") + .create_async() + .await; + + let adapter = make_runway(&server); + let req = VideoRequest::builder("gen-3", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ video_poll 正常路径 ============ + + #[tokio::test] + async fn runway_video_poll_success_returns_video_url() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "asset-abc", + "status": "completed", + "url": "https://example.com/video.mp4", + "progress": 100, + "createdAt": 1700000000, + "updatedAt": 1700000100 + }); + let mock = server + .mock("GET", "/assets/asset-abc") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter + .video_poll("asset-abc", "gen-3") + .await + .expect("video_poll 应成功"); + + assert_eq!(status.task_id, "asset-abc"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/video.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert_eq!(status.created_at, Some(1700000000)); + assert_eq!(status.updated_at, Some(1700000100)); + assert!(status.error.is_none()); + mock.assert_async().await; + } + + #[tokio::test] + async fn runway_video_poll_processing_returns_progress() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-1") + .with_status(200) + .with_body(json!({"id": "asset-1", "status": "processing", "progress": 45}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-1", "gen-3").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(45)); + assert!(status.video_url.is_none()); + } + + #[tokio::test] + async fn runway_video_poll_failed_returns_error() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-2") + .with_status(200) + .with_body( + json!({"id": "asset-2", "status": "failed", "error": "content policy violation"}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-2", "gen-3").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("content policy violation")); + } + + #[tokio::test] + async fn runway_video_poll_failed_returns_error_message_field() { + // errorMessage 字段兜底 + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-3") + .with_status(200) + .with_body( + json!({"id": "asset-3", "status": "error", "errorMessage": "internal failure"}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-3", "gen-3").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("internal failure")); + } + + #[tokio::test] + async fn runway_video_poll_queued_maps_to_pending() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-4") + .with_status(200) + .with_body(json!({"id": "asset-4", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-4", "gen-3").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn runway_video_poll_video_url_output_path() { + // output.url 兜底路径 + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-5") + .with_status(200) + .with_body( + json!({ + "id": "asset-5", + "status": "succeeded", + "output": {"url": "https://example.com/out.mp4"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-5", "gen-3").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/out.mp4") + ); + } + + #[tokio::test] + async fn runway_video_poll_video_url_assets_array_path() { + // assets[0].url 兜底路径 + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-6") + .with_status(200) + .with_body( + json!({ + "id": "asset-6", + "status": "completed", + "assets": [{"url": "https://example.com/arr.mp4"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-6", "gen-3").await.unwrap(); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/arr.mp4") + ); + } + + #[tokio::test] + async fn runway_video_poll_video_url_videourl_field() { + // videoUrl 字段兜底 + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-7") + .with_status(200) + .with_body( + json!({ + "id": "asset-7", + "status": "completed", + "videoUrl": "https://example.com/vu.mp4" + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_runway(&server); + let status = adapter.video_poll("asset-7", "gen-3").await.unwrap(); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/vu.mp4") + ); + } + + // ============ video_poll 错误路径 ============ + + #[tokio::test] + async fn runway_video_poll_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-x") + .with_status(401) + .with_body(json!({"error": "bad key"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let err = adapter.video_poll("asset-x", "gen-3").await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn runway_video_poll_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/nonexistent") + .with_status(404) + .with_body(json!({"error": "asset not found"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let err = adapter + .video_poll("nonexistent", "gen-3") + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn runway_video_poll_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("GET", "/assets/asset-y") + .with_status(429) + .with_body(json!({"error": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_runway(&server); + let err = adapter.video_poll("asset-y", "gen-3").await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ============ 不支持的能力 ============ + + #[tokio::test] + async fn runway_chat_returns_unsupported() { + let adapter = make_runway_no_server(); + let req = ChatRequest::builder("gen-3", vec![ChatMessage::user("hi")]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn runway_image_generate_returns_unsupported() { + let adapter = make_runway_no_server(); + let req = ImageRequest::builder("gen-3", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn runway_chat_stream_returns_unsupported() { + let adapter = make_runway_no_server(); + let req = ChatRequest::builder("gen-3", vec![ChatMessage::user("hi")]).build(); + let result = adapter.chat_stream(req).await; + assert!(matches!( + result, + Err(AibridgeError::UnsupportedCapability { .. }) + )); + } + + // ============ list_models ============ + + #[tokio::test] + async fn runway_list_models_returns_hardcoded() { + let adapter = make_runway_no_server(); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "gen-3"); + assert_eq!(models[0].name, "Gen-3 Alpha"); + assert_eq!(models[0].provider, "runway"); + assert_eq!(models[0].model_type, ModelType::Video); + assert_eq!(models[1].id, "gen-3-turbo"); + } + + #[tokio::test] + async fn runway_list_models_filter_by_video_type() { + let adapter = make_runway_no_server(); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert_eq!(videos.len(), 2); + assert!(videos.iter().all(|m| m.model_type == ModelType::Video)); + } + + #[tokio::test] + async fn runway_list_models_filter_by_image_returns_empty() { + let adapter = make_runway_no_server(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert!(images.is_empty()); + } + + // ============ start / close ============ + + #[tokio::test] + async fn runway_start_and_close_are_noops() { + let mut adapter = make_runway_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn runway_map_api_error_401_is_authentication() { + let err = RunwayAdapter::map_api_error(401, ""); + match err { + AibridgeError::Authentication { message } => assert!(message.contains("Runway")), + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn runway_map_api_error_403_is_authentication() { + let err = RunwayAdapter::map_api_error(403, ""); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn runway_map_api_error_429_is_rate_limit() { + let err = RunwayAdapter::map_api_error(429, ""); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn runway_map_api_error_404_is_model_not_found() { + let err = RunwayAdapter::map_api_error(404, ""); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[test] + fn runway_map_api_error_400_is_validation() { + let err = RunwayAdapter::map_api_error(400, ""); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn runway_map_api_error_400_carries_details() { + let body = json!({"error": "bad param"}).to_string(); + let err = RunwayAdapter::map_api_error(400, &body); + match err { + AibridgeError::Validation { message, details } => { + assert_eq!(message, "bad param"); + assert_eq!(details["error"], "bad param"); + } + _ => panic!("应为 Validation"), + } + } + + #[test] + fn runway_map_api_error_500_is_api() { + let err = RunwayAdapter::map_api_error(500, ""); + match err { + AibridgeError::Api { status, .. } => assert_eq!(status, 500), + _ => panic!("应为 Api"), + } + } + + #[test] + fn runway_map_api_error_extracts_error_message() { + let body = json!({"error": "something went wrong"}).to_string(); + let err = RunwayAdapter::map_api_error(500, &body); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "something went wrong"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn runway_map_api_error_extracts_openai_style_message() { + let body = json!({"error": {"message": "rate limited"}}).to_string(); + let err = RunwayAdapter::map_api_error(429, &body); + match err { + AibridgeError::RateLimit { message, .. } => { + // 429 固定消息,不读 body + assert!(message.contains("rate limit")); + } + _ => panic!("应为 RateLimit"), + } + } + + #[test] + fn runway_map_api_error_no_json_falls_back_to_http_status() { + let err = RunwayAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ map_runway_status 单元测试 ============ + + #[test] + fn map_runway_status_pending_is_pending() { + assert_eq!(map_runway_status("pending"), TaskStatus::Pending); + } + + #[test] + fn map_runway_status_queued_is_pending() { + assert_eq!(map_runway_status("queued"), TaskStatus::Pending); + } + + #[test] + fn map_runway_status_processing_is_processing() { + assert_eq!(map_runway_status("processing"), TaskStatus::Processing); + assert_eq!(map_runway_status("running"), TaskStatus::Processing); + assert_eq!(map_runway_status("in_progress"), TaskStatus::Processing); + } + + #[test] + fn map_runway_status_success_variants() { + assert_eq!(map_runway_status("completed"), TaskStatus::Success); + assert_eq!(map_runway_status("succeeded"), TaskStatus::Success); + assert_eq!(map_runway_status("success"), TaskStatus::Success); + } + + #[test] + fn map_runway_status_failed_variants() { + assert_eq!(map_runway_status("failed"), TaskStatus::Failed); + assert_eq!(map_runway_status("error"), TaskStatus::Failed); + assert_eq!(map_runway_status("cancelled"), TaskStatus::Failed); + } + + #[test] + fn map_runway_status_case_insensitive() { + assert_eq!(map_runway_status("PENDING"), TaskStatus::Pending); + assert_eq!(map_runway_status("Completed"), TaskStatus::Success); + assert_eq!(map_runway_status("FAILED"), TaskStatus::Failed); + } + + #[test] + fn map_runway_status_unknown_defaults_to_pending() { + assert_eq!(map_runway_status("unknown_state"), TaskStatus::Pending); + assert_eq!(map_runway_status(""), TaskStatus::Pending); + } + + // ============ file_input_to_url 单元测试 ============ + + #[test] + fn file_input_to_url_returns_url_for_url_variant() { + let f = FileInput::url("https://example.com/x.png"); + assert_eq!(file_input_to_url(&f), "https://example.com/x.png"); + } + + #[test] + fn file_input_to_url_returns_base64_for_base64_variant() { + let f = FileInput::base64("aGVsbG8="); + assert_eq!(file_input_to_url(&f), "aGVsbG8="); + } + + #[test] + fn file_input_to_url_returns_empty_for_path_variant() { + let f = FileInput::path("/tmp/x.png"); + assert_eq!(file_input_to_url(&f), ""); + } + + #[test] + fn file_input_to_url_returns_empty_for_bytes_variant() { + let f = FileInput::bytes(vec![1, 2, 3]); + assert_eq!(file_input_to_url(&f), ""); + } + + // ============ build_create_body 单元测试 ============ + + #[test] + fn build_create_body_text2video_includes_fields() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("runway", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = RunwayAdapter::with_http(http, config); + let req = VideoRequest::builder("gen-3", "a cat") + .aspect_ratio("16:9") + .seed(42) + .build(); + let (endpoint, body) = adapter.build_create_body(&req); + assert_eq!(endpoint, "text_to_video"); + assert_eq!(body["model"], "gen-3"); + assert_eq!(body["prompt"], "a cat"); + assert_eq!(body["aspectRatio"], "16:9"); + assert_eq!(body["seed"], 42); + assert_eq!(body["width"], 1280); + assert_eq!(body["height"], 720); + } + + #[test] + fn build_create_body_image2video_uses_prompt_image() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("runway", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = RunwayAdapter::with_http(http, config); + let req = VideoRequest::builder("gen-3", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let (endpoint, body) = adapter.build_create_body(&req); + assert_eq!(endpoint, "image_to_video"); + assert_eq!(body["promptImage"], "https://example.com/a.png"); + assert_eq!(body["promptText"], "animate"); + // text_to_video 的 prompt 字段不应出现 + assert!(body.get("prompt").is_none()); + } + + #[test] + fn build_create_body_merges_extra() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("runway", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = RunwayAdapter::with_http(http, config); + let req = VideoRequest::builder("gen-3", "a cat") + .extra("custom", "value") + .build(); + let (_, body) = adapter.build_create_body(&req); + assert_eq!(body["custom"], "value"); + } + + #[test] + fn build_create_body_uses_default_model_when_empty() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("runway", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = RunwayAdapter::with_http(http, config); + let req = VideoRequest::builder("", "a cat").build(); + let (_, body) = adapter.build_create_body(&req); + assert_eq!(body["model"], "gen-3"); + } +} From 0124c5c0e0d78cfa6338640f87ae8d2d436ae227 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 00:57:01 +0800 Subject: [PATCH 37/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20pika=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/pika.rs | 1594 +++++++++++++++++++++ 1 file changed, 1594 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/pika.rs diff --git a/crates/aibridge-core/src/adapters/pika.rs b/crates/aibridge-core/src/adapters/pika.rs new file mode 100644 index 0000000..502c2ab --- /dev/null +++ b/crates/aibridge-core/src/adapters/pika.rs @@ -0,0 +1,1594 @@ +//! Pika 适配器(视频生成,独立协议) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/pika.py`。 +//! +//! Pika API 为**独立协议**(非 OpenAI 兼容),不复用 `OpenAiCompatAdapter`: +//! - 创建视频任务:`POST /generations`(请求体用 `prompt_text` 字段,图生视频用 `prompt_image`) +//! - 查询任务状态:`GET /generations/{task_id}` +//! - 认证:`Bearer ` +//! - 默认 Base URL:`https://api.pika.art/v1` +//! +//! ## 阶段范围 +//! 阶段 2b 实现 video_create + video_poll + list_models(硬编码)。 +//! 不支持的能力(chat / image / embed / audio)走 `Adapter` trait 默认实现返 +//! `UnsupportedCapability`,与 Python v1 抛 `UnsupportedCapabilityError` 行为一致。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::common::{ModelInfo, ModelType, TaskStatus, VideoMode}; +use crate::model::image::FileInput; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +/// Pika 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +pub const DEFAULT_PIKA_BASE_URL: &str = "https://api.pika.art/v1"; + +/// 默认视频宽度(与 `VideoRequest` 默认值一致,用于判断是否需要发送 width 字段) +const DEFAULT_WIDTH: u32 = 1280; +/// 默认视频高度(与 `VideoRequest` 默认值一致,用于判断是否需要发送 height 字段) +const DEFAULT_HEIGHT: u32 = 720; + +// ==================== 能力集合构造 ==================== + +/// Pika 支持的能力集合 +/// +/// 对齐 Python v1 `PikaAdapter.supported_capabilities = ["video"]`。 +/// Rust 用 `VideoGenerate` 表达视频生成能力(含 video_create + video_poll), +/// 并声明 text2video / image2video 两个子能力。 +fn pika_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps +} + +// ==================== Pika 视频生成适配器 ==================== + +/// Pika 适配器 +/// +/// Pika 视频生成平台,支持文生视频与图生视频(Pika 1.0 / 2.0)。 +/// 官方 API 文档:https://pika.art/ +/// +/// ## API 规范 +/// - Base URL: `https://api.pika.art/v1` +/// - 创建生成: `POST /generations` +/// - 查询状态: `GET /generations/{id}` +/// - 认证: `Authorization: Bearer {api_key}` +/// +/// ## 阶段范围 +/// 阶段 2b 实现 video_create + video_poll + list_models(硬编码)。 +pub struct PikaAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl PikaAdapter { + /// 创建 Pika 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_PIKA_BASE_URL`] 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_PIKA_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: pika_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_PIKA_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: pika_capabilities(), + } + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: pika)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用 Pika 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 Bearer 认证的 GET 请求,并用 Pika 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 Pika `/generations` 请求体 + /// + /// 移植自 Python v1 `video_create`: + /// - `model` / `prompt_text` 必选 + /// - image2video 模式:`reference_images[0]` → `prompt_image`(优先),`first_frame` 次之 + /// - `width` / `height` 仅在非默认值时发送(对齐 Python「显式传才发送」语义) + /// - `seed` / `aspect_ratio` / `duration` / `negative_prompt_text` 可选透传 + /// - `extra` 字段合并到顶层(透传厂商特有参数) + fn build_generations_body(&self, req: &VideoRequest) -> Value { + let mut body = json!({ + "model": req.model, + "prompt_text": req.prompt, + }); + + // image2video 模式:reference_images[0] 优先作为 prompt_image,first_frame 次之 + let is_image2video = matches!(req.mode, VideoMode::Image2Video); + if is_image2video { + if let Some(url) = req.reference_images.first().and_then(file_input_to_url) { + body["prompt_image"] = json!(url); + } else if let Some(ff) = req.first_frame.as_ref().and_then(file_input_to_url) { + body["prompt_image"] = json!(ff); + } + } + + // width / height 仅在非默认值时发送(默认 1280x720 不发送,让 Pika 用自身默认或 aspect_ratio) + if req.width != DEFAULT_WIDTH { + body["width"] = json!(req.width); + } + if req.height != DEFAULT_HEIGHT { + body["height"] = json!(req.height); + } + if let Some(seed) = req.seed { + body["seed"] = json!(seed); + } + if let Some(ar) = &req.aspect_ratio { + body["aspect_ratio"] = json!(ar); + } + if let Some(d) = req.duration { + body["duration"] = json!(d); + } + if let Some(np) = &req.negative_prompt { + body["negative_prompt_text"] = json!(np); + } + + // extra 透传(合并到顶层) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 解析 Pika `/generations` 创建响应 → VideoTask + /// + /// Pika 创建响应任务 ID 可能在 `id` / `generation_id` / `taskId` 字段, + /// 均缺失时回退到生成的 `vid_` 前缀 ID(与 Python v1 一致)。 + fn parse_video_task(value: &Value, model: &str) -> Result { + let task_id = value + .get("id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("generation_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("taskId") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .unwrap_or_else(|| util::generate_id("vid")); + let raw_status = value + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("pending"); + let created_at = value + .get("created_at") + .and_then(|v| v.as_u64()) + .unwrap_or_else(util::current_timestamp); + Ok(VideoTask { + task_id, + model: model.to_string(), + status: map_pika_status(raw_status), + created_at, + }) + } + + /// 解析 Pika `/generations/{id}` 查询响应 → VideoStatus + /// + /// 移植自 Python v1 `video_poll`: + /// - 视频 URL 兼容多路径:`video_url` / `url` / `output.video_url` / `output.url` / `results[0].url` + /// - 错误信息兼容:`error` / `error_message` / `failure_reason` + /// - `progress` / `created_at` / `updated_at` 透传(数字形式) + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let raw_status = value + .get("status") + .and_then(|v| v.as_str()) + .unwrap_or("pending"); + let status = map_pika_status(raw_status); + + // 视频 URL:多路径兼容(无条件提取,与 Python v1 一致) + let video_url = value + .get("video_url") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| value.get("url").and_then(|v| v.as_str()).map(str::to_owned)) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("video_url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("output") + .and_then(|o| o.get("url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("results") + .and_then(|r| r.as_array()) + .and_then(|arr| arr.first()) + .and_then(|f| f.get("url")) + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + + // 错误信息:多字段兼容(无条件提取,与 Python v1 一致) + let error = value + .get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("error_message") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .or_else(|| { + value + .get("failure_reason") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + + let progress = value + .get("progress") + .and_then(|v| v.as_u64()) + .map(|p| p as u32); + let created_at = value.get("created_at").and_then(|v| v.as_u64()); + let updated_at = value.get("updated_at").and_then(|v| v.as_u64()); + + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress, + error, + created_at, + updated_at, + } + } + + /// 将 Pika API 错误响应映射为 AibridgeError + /// + /// 移植自 Python v1 `_handle_pika_error` 并按阶段 2b 统一错误映射要求调整: + /// - 401 → Authentication("Invalid Pika API key") + /// - 429 → RateLimit("Pika rate limit exceeded") + /// - 404 → ModelNotFound(Pika 资源/任务不存在) + /// - 400 → Validation(请求参数校验错误) + /// - 其余 ≥400 → Api(提取 error.message / message / error / detail,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Pika API key".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Pika rate limit exceeded".to_string(), + retry_after: None, + }, + 404 => AibridgeError::ModelNotFound { + model: "Pika generation or resource not found".to_string(), + }, + 400 => AibridgeError::validation_with_details( + parse_error_message(body, status), + serde_json::from_str(body).unwrap_or(serde_json::Value::Null), + ), + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for PikaAdapter { + fn provider_type(&self) -> &str { + "pika" + } + + fn provider_name(&self) -> &str { + "Pika" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 创建视频生成任务:`POST /generations` + async fn video_create(&self, req: VideoRequest) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let body = self.build_generations_body(&req); + let value = self.post_authed_json("generations", &body).await?; + Self::parse_video_task(&value, &req.model) + } + + /// 查询视频任务状态:`GET /generations/{task_id}` + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let path = format!("generations/{task_id}"); + let value = self.get_authed_json(&path).await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 模型列表(硬编码) + /// + /// Pika 无标准 `/models` 端点,暂保留硬编码列表(与 Python v1 一致)。 + /// 含 pika-1.0 / pika-2 两个模型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = pika_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / image / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 内部:辅助函数 ==================== + +/// 映射 Pika 状态字符串到统一 TaskStatus +/// +/// 移植自 Python v1 `_map_pika_status`: +/// - pending / queued / in_queue → Pending +/// - processing / in_progress / generating → Processing +/// - completed / finished / succeeded / success → Success +/// - failed / error / failure / cancelled → Failed +/// - 未知 → Pending(与 Python 默认值一致) +fn map_pika_status(raw: &str) -> TaskStatus { + match raw.to_lowercase().as_str() { + "pending" | "queued" | "in_queue" => TaskStatus::Pending, + "processing" | "in_progress" | "generating" => TaskStatus::Processing, + "completed" | "finished" | "succeeded" | "success" => TaskStatus::Success, + "failed" | "error" | "failure" | "cancelled" => TaskStatus::Failed, + _ => TaskStatus::Pending, + } +} + +/// 从 FileInput 提取 URL 字符串 +/// +/// Pika 的 `prompt_image` 字段接受 URL 或 base64 字符串: +/// - `Url(s)` / `Base64(s)` → `Some(s)` +/// - `Path(_)` / `Bytes(_)` → `None`(需调用方先上传为可访问 URL) +fn file_input_to_url(input: &FileInput) -> Option { + match input { + FileInput::Url(s) | FileInput::Base64(s) => Some(s.clone()), + FileInput::Path(_) | FileInput::Bytes(_) => None, + } +} + +/// 解析错误体中的 message 字段 +/// +/// 通用错误消息提取,兼容多种错误体结构: +/// - `{"error": {"message": "..."}}`(OpenAI 风格) +/// - `{"error": "..."}`(顶层 error 字符串) +/// - `{"message": "..."}`(顶层 message) +/// - `{"detail": "..."}`(顶层 detail) +/// +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + // error.message(OpenAI 风格) + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 顶层 error 字符串 + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 message + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 detail + if let Some(msg) = v.get("detail").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +/// Pika 硬编码模型列表 +/// +/// 对应 Python v1 `PikaAdapter.list_models`。 +/// 注意:该 Provider 无标准 `/models` 端点,暂保留硬编码列表。 +fn pika_hardcoded_models() -> Vec { + vec![ + ModelInfo { + id: "pika-1.0".into(), + name: "Pika 1.0".into(), + model_type: ModelType::Video, + provider: "pika".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Pika 1.0 视频生成模型".into()), + created: None, + }, + ModelInfo { + id: "pika-2".into(), + name: "Pika 2.0".into(), + model_type: ModelType::Video, + provider: "pika".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Pika 2.0 最新视频生成模型".into()), + created: None, + }, + ] +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::chat::ChatRequest; + use crate::model::image::ImageRequest; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 PikaAdapter(指向 mockito server) + fn make_pika(server: &Server) -> PikaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("pika", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + PikaAdapter::with_http(http, config) + } + + /// 构造不指向任何 server 的 PikaAdapter(用于不发请求的元信息/能力测试) + fn make_pika_no_server() -> PikaAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_PIKA_BASE_URL) + .build(); + let config = ProviderConfig::from_options("pika", opts); + PikaAdapter::new(config).expect("PikaAdapter 构造应成功") + } + + // ============ 元信息 ============ + + #[test] + fn pika_provider_type_and_name_match_python() { + let adapter = make_pika_no_server(); + assert_eq!(adapter.provider_type(), "pika"); + assert_eq!(adapter.provider_name(), "Pika"); + } + + #[test] + fn pika_requires_api_key_is_true() { + let adapter = make_pika_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn pika_capabilities_contains_only_video() { + let adapter = make_pika_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + // chat / image 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn pika_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("pika", opts); + let adapter = PikaAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_PIKA_BASE_URL); + } + + #[test] + fn pika_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.pika-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("pika", opts); + let adapter = PikaAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.pika-proxy.com/v1"); + } + + #[test] + fn pika_base_url_ignores_empty_string() { + // 空白 base_url 应回退到默认值 + let opts = ClientOptions::builder() + .api_key("k") + .base_url(" ") + .build(); + let config = ProviderConfig::from_options("pika", opts); + let adapter = PikaAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_PIKA_BASE_URL); + } + + // ============ video_create 正常路径 ============ + + #[tokio::test] + async fn pika_video_create_success_returns_task() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "status": "pending", + "created_at": 1700000000 + }); + let mock = server + .mock("POST", "/generations") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat running") + .aspect_ratio("16:9") + .duration(5) + .build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + + assert_eq!(task.task_id, "gen-abc"); + assert_eq!(task.model, "pika-1.0"); + assert_eq!(task.status, TaskStatus::Pending); + assert_eq!(task.created_at, 1700000000); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_sends_model_and_prompt_text() { + // 验证请求体用 prompt_text 字段(非 prompt),且默认不发送 width/height + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "model": "pika-1.0", + "prompt_text": "a cat" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_passes_optional_params() { + // aspect_ratio / duration / seed / negative_prompt 透传 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "aspect_ratio": "16:9", + "duration": 5, + "seed": 42, + "negative_prompt_text": "blurry" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat") + .aspect_ratio("16:9") + .duration(5) + .seed(42) + .negative_prompt("blurry") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_passes_width_height_when_non_default() { + // 非默认 width/height 才发送 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "width": 1920, + "height": 1080 + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat") + .width(1920) + .height(1080) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_default_width_height_not_sent() { + // 默认 1280x720 不应出现在请求体 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::Json(json!({ + "model": "pika-1.0", + "prompt_text": "a cat" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_image2video_with_reference_image() { + // image2video 模式:reference_images[0] → prompt_image + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "prompt_image": "https://example.com/a.png" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "animate this") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_image2video_with_first_frame() { + // reference_images 为空时,first_frame 作为 prompt_image 后备 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "prompt_image": "https://example.com/first.png" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "animate this") + .mode(VideoMode::Image2Video) + .first_frame(FileInput::url("https://example.com/first.png")) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_text2video_does_not_send_prompt_image() { + // text2video 模式即使有 reference_images 也不发送 prompt_image + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::Json(json!({ + "model": "pika-1.0", + "prompt_text": "a cat" + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat") + .mode(VideoMode::Text2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_passes_extra_params() { + // extra 字段透传到顶层 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "custom_param": "custom_value", + "options": {"frame_rate": 30} + }))) + .with_status(200) + .with_body(json!({"id": "x", "status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat") + .extra("custom_param", "custom_value") + .extra("options", json!({"frame_rate": 30})) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_create_accepts_generation_id_field() { + // 任务 ID 在 generation_id 字段 + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(200) + .with_body(json!({"generation_id": "gen-xyz", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "gen-xyz"); + assert_eq!(task.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn pika_video_create_accepts_task_id_field() { + // 任务 ID 在 taskId 字段 + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(200) + .with_body(json!({"taskId": "task-001", "status": "processing"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "task-001"); + assert_eq!(task.status, TaskStatus::Processing); + } + + #[tokio::test] + async fn pika_video_create_uses_generated_id_when_missing() { + // 响应缺任务 ID 时,回退到生成的 vid_ ID + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(200) + .with_body(json!({"status": "pending"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert!(task.task_id.starts_with("vid_")); + } + + // ============ video_create 错误路径 ============ + + #[tokio::test] + async fn pika_video_create_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(401) + .with_body(json!({"error": {"message": "invalid key"}}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn pika_video_create_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(429) + .with_body(json!({"message": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn pika_video_create_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(404) + .with_body(json!({"error": "not found"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn pika_video_create_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(400) + .with_body(json!({"error": {"message": "invalid prompt"}}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("invalid prompt")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn pika_video_create_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert!(message.contains("internal")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn pika_video_create_error_no_json_body_falls_back() { + let mut server = Server::new_async().await; + server + .mock("POST", "/generations") + .with_status(502) + .with_body("Bad Gateway") + .create_async() + .await; + + let adapter = make_pika(&server); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ video_poll 正常路径 ============ + + #[tokio::test] + async fn pika_video_poll_success_returns_video_url() { + let mut server = Server::new_async().await; + let body = json!({ + "id": "gen-abc", + "status": "completed", + "video_url": "https://example.com/video.mp4", + "progress": 100, + "created_at": 1700000000, + "updated_at": 1700000100 + }); + let mock = server + .mock("GET", "/generations/gen-abc") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter + .video_poll("gen-abc", "pika-1.0") + .await + .expect("video_poll 应成功"); + + assert_eq!(status.task_id, "gen-abc"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/video.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert_eq!(status.created_at, Some(1700000000)); + assert_eq!(status.updated_at, Some(1700000100)); + assert!(status.error.is_none()); + mock.assert_async().await; + } + + #[tokio::test] + async fn pika_video_poll_processing_status() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-1") + .with_status(200) + .with_body( + json!({ + "id": "gen-1", + "status": "processing", + "progress": 45 + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-1", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(45)); + } + + #[tokio::test] + async fn pika_video_poll_failed_returns_error() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-2") + .with_status(200) + .with_body( + json!({ + "id": "gen-2", + "status": "failed", + "failure_reason": "content policy violation" + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-2", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("content policy violation")); + } + + #[tokio::test] + async fn pika_video_poll_failed_extracts_error_message_field() { + // error_message 字段 + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-3") + .with_status(200) + .with_body( + json!({ + "id": "gen-3", + "status": "error", + "error_message": "internal error" + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-3", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("internal error")); + } + + #[tokio::test] + async fn pika_video_poll_compatible_url_paths() { + // 兼容 output.video_url / output.url / results[0].url / 顶层 url + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-4") + .with_status(200) + .with_body( + json!({ + "id": "gen-4", + "status": "finished", + "output": {"video_url": "https://legacy.example.com/v.mp4"} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-4", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://legacy.example.com/v.mp4") + ); + } + + #[tokio::test] + async fn pika_video_poll_url_from_results_array() { + // results[0].url 路径 + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-5") + .with_status(200) + .with_body( + json!({ + "id": "gen-5", + "status": "succeeded", + "results": [{"url": "https://results.example.com/v.mp4"}] + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-5", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://results.example.com/v.mp4") + ); + } + + #[tokio::test] + async fn pika_video_poll_queued_status_maps_to_pending() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-6") + .with_status(200) + .with_body(json!({"id": "gen-6", "status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-6", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn pika_video_poll_cancelled_maps_to_failed() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-7") + .with_status(200) + .with_body( + json!({"id": "gen-7", "status": "cancelled", "error": "user cancelled"}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_pika(&server); + let status = adapter.video_poll("gen-7", "pika-1.0").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("user cancelled")); + } + + // ============ video_poll 错误路径 ============ + + #[tokio::test] + async fn pika_video_poll_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-x") + .with_status(401) + .with_body(json!({"error": {"message": "bad key"}}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let err = adapter.video_poll("gen-x", "pika-1.0").await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn pika_video_poll_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/nonexistent") + .with_status(404) + .with_body(json!({"error": "not found"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let err = adapter + .video_poll("nonexistent", "pika-1.0") + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn pika_video_poll_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("GET", "/generations/gen-y") + .with_status(429) + .with_body(json!({"message": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_pika(&server); + let err = adapter.video_poll("gen-y", "pika-1.0").await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + // ============ 不支持的能力 ============ + + #[tokio::test] + async fn pika_chat_returns_unsupported() { + let adapter = make_pika_no_server(); + let req = ChatRequest::builder("pika-1.0", vec![]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn pika_image_generate_returns_unsupported() { + let adapter = make_pika_no_server(); + let req = ImageRequest::builder("pika-1.0", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn pika_embed_returns_unsupported() { + let adapter = make_pika_no_server(); + let req = crate::model::options::EmbedRequest { + model: "pika-1.0".into(), + input: crate::model::options::EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: std::collections::HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ list_models ============ + + #[tokio::test] + async fn pika_list_models_returns_hardcoded() { + let adapter = make_pika_no_server(); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + assert_eq!(models[0].id, "pika-1.0"); + assert_eq!(models[0].provider, "pika"); + assert_eq!(models[0].model_type, ModelType::Video); + assert_eq!(models[1].id, "pika-2"); + } + + #[tokio::test] + async fn pika_list_models_filter_by_video_type() { + let adapter = make_pika_no_server(); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert_eq!(videos.len(), 2); + assert!(videos.iter().all(|m| m.model_type == ModelType::Video)); + } + + #[tokio::test] + async fn pika_list_models_filter_by_image_returns_empty() { + let adapter = make_pika_no_server(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert!(images.is_empty()); + } + + // ============ start / close ============ + + #[tokio::test] + async fn pika_start_and_close_are_noops() { + let mut adapter = make_pika_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn map_api_error_401_is_authentication() { + let err = PikaAdapter::map_api_error(401, "{\"error\":{\"message\":\"bad\"}}"); + match err { + AibridgeError::Authentication { message } => assert!(message.contains("Pika")), + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn map_api_error_429_is_rate_limit() { + let err = PikaAdapter::map_api_error(429, "{}"); + match err { + AibridgeError::RateLimit { + message, + retry_after, + } => { + assert!(message.contains("Pika")); + assert!(retry_after.is_none()); + } + _ => panic!("应为 RateLimit"), + } + } + + #[test] + fn map_api_error_404_is_model_not_found() { + let err = PikaAdapter::map_api_error(404, "{}"); + match err { + AibridgeError::ModelNotFound { model } => assert!(model.contains("Pika")), + _ => panic!("应为 ModelNotFound"), + } + } + + #[test] + fn map_api_error_400_is_validation() { + let err = PikaAdapter::map_api_error(400, "{\"error\":{\"message\":\"bad param\"}}"); + match err { + AibridgeError::Validation { message, .. } => assert_eq!(message, "bad param"), + _ => panic!("应为 Validation"), + } + } + + #[test] + fn map_api_error_500_extracts_message() { + let err = PikaAdapter::map_api_error(500, "{\"error\":{\"message\":\"internal\"}}"); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_extracts_top_level_message() { + let err = PikaAdapter::map_api_error(503, "{\"message\":\"unavailable\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "unavailable"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_extracts_detail_field() { + let err = PikaAdapter::map_api_error(422, "{\"detail\":\"validation failed\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "validation failed"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_falls_back_to_http_status() { + let err = PikaAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + // ============ map_pika_status 单元测试 ============ + + #[test] + fn map_pika_status_pending_variants() { + assert_eq!(map_pika_status("pending"), TaskStatus::Pending); + assert_eq!(map_pika_status("queued"), TaskStatus::Pending); + assert_eq!(map_pika_status("in_queue"), TaskStatus::Pending); + } + + #[test] + fn map_pika_status_processing_variants() { + assert_eq!(map_pika_status("processing"), TaskStatus::Processing); + assert_eq!(map_pika_status("in_progress"), TaskStatus::Processing); + assert_eq!(map_pika_status("generating"), TaskStatus::Processing); + } + + #[test] + fn map_pika_status_success_variants() { + assert_eq!(map_pika_status("completed"), TaskStatus::Success); + assert_eq!(map_pika_status("finished"), TaskStatus::Success); + assert_eq!(map_pika_status("succeeded"), TaskStatus::Success); + assert_eq!(map_pika_status("success"), TaskStatus::Success); + } + + #[test] + fn map_pika_status_failed_variants() { + assert_eq!(map_pika_status("failed"), TaskStatus::Failed); + assert_eq!(map_pika_status("error"), TaskStatus::Failed); + assert_eq!(map_pika_status("failure"), TaskStatus::Failed); + assert_eq!(map_pika_status("cancelled"), TaskStatus::Failed); + } + + #[test] + fn map_pika_status_case_insensitive() { + assert_eq!(map_pika_status("PENDING"), TaskStatus::Pending); + assert_eq!(map_pika_status("Completed"), TaskStatus::Success); + assert_eq!(map_pika_status("FAILED"), TaskStatus::Failed); + } + + #[test] + fn map_pika_status_unknown_defaults_to_pending() { + assert_eq!(map_pika_status("unknown_state"), TaskStatus::Pending); + assert_eq!(map_pika_status(""), TaskStatus::Pending); + } + + // ============ file_input_to_url 单元测试 ============ + + #[test] + fn file_input_to_url_returns_url_for_url_variant() { + let f = FileInput::url("https://example.com/x.png"); + assert_eq!( + file_input_to_url(&f), + Some("https://example.com/x.png".to_string()) + ); + } + + #[test] + fn file_input_to_url_returns_base64_for_base64_variant() { + let f = FileInput::base64("aGVsbG8="); + assert_eq!(file_input_to_url(&f), Some("aGVsbG8=".to_string())); + } + + #[test] + fn file_input_to_url_returns_none_for_path_and_bytes() { + assert_eq!(file_input_to_url(&FileInput::path("/tmp/x")), None); + assert_eq!(file_input_to_url(&FileInput::bytes(vec![1, 2])), None); + } + + // ============ build_generations_body 单元测试 ============ + + #[test] + fn build_generations_body_includes_model_and_prompt_text() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("pika", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = PikaAdapter::with_http(http, config); + let req = VideoRequest::builder("pika-1.0", "a cat").build(); + let body = adapter.build_generations_body(&req); + assert_eq!(body["model"], "pika-1.0"); + assert_eq!(body["prompt_text"], "a cat"); + // 默认 width/height 不发送 + assert!(body.get("width").is_none()); + assert!(body.get("height").is_none()); + // 无 prompt_image + assert!(body.get("prompt_image").is_none()); + } + + #[test] + fn build_generations_body_image2video_sets_prompt_image() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("pika", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = PikaAdapter::with_http(http, config); + let req = VideoRequest::builder("pika-1.0", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let body = adapter.build_generations_body(&req); + assert_eq!(body["prompt_image"], "https://example.com/a.png"); + } + + #[test] + fn build_generations_body_passes_optional_fields() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("pika", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = PikaAdapter::with_http(http, config); + let req = VideoRequest::builder("pika-2", "a cat") + .width(1920) + .height(1080) + .aspect_ratio("16:9") + .duration(5) + .seed(42) + .negative_prompt("blurry") + .build(); + let body = adapter.build_generations_body(&req); + assert_eq!(body["width"], 1920); + assert_eq!(body["height"], 1080); + assert_eq!(body["aspect_ratio"], "16:9"); + assert_eq!(body["duration"], 5); + assert_eq!(body["seed"], 42); + assert_eq!(body["negative_prompt_text"], "blurry"); + } + + #[test] + fn build_generations_body_extra_passthrough() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("pika", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = PikaAdapter::with_http(http, config); + let req = VideoRequest::builder("pika-1.0", "a cat") + .extra("custom_param", "custom_value") + .build(); + let body = adapter.build_generations_body(&req); + assert_eq!(body["custom_param"], "custom_value"); + } +} From 3f05c391fcd86b00dbea3d9884a4a151c8239fa0 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:03:12 +0800 Subject: [PATCH 38/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20=E7=AC=AC=E4=BA=8C=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20runway=20+=20pika=20=E5=88=B0=E5=B7=A5=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 38 +++++++++++++++++---- crates/aibridge-core/src/adapters/mod.rs | 6 ++++ 2 files changed, 37 insertions(+), 7 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index e905ff6..c902c47 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -27,6 +27,8 @@ use crate::adapters::more_models::{ CohereAdapter, DeepSeekAdapter, MistralAdapter, PerplexityAdapter, StepFunAdapter, }; use crate::adapters::openai::OpenAiAdapter; +use crate::adapters::pika::PikaAdapter; +use crate::adapters::runway::RunwayAdapter; use crate::adapters::stability::StabilityAdapter; use crate::adapters::volcengine_cv::VolcengineCvAdapter; use crate::config::ProviderConfig; @@ -72,9 +74,10 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ // 阶段 2b 第一批独立协议(已实现): "anthropic", "stability", - // 阶段 2b/2c 待实现: + // 阶段 2b 第二批独立协议(已实现): "runway", "pika", + // 阶段 2b/2c 待实现: "kling", "edge-tts", "elevenlabs", @@ -134,11 +137,15 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { // 阶段 2b 独立协议:别名对齐 Python agn/adapters/{anthropic,stability}.py 末尾 register 调用(均无别名) "anthropic" => Ok(Box::new(AnthropicAdapter::new(config)?)), "stability" => Ok(Box::new(StabilityAdapter::new(config)?)), + // 阶段 2b 第二批独立协议:别名对齐 Python agn/adapters/{runway,pika}.py 末尾 register 调用(均无别名) + "runway" => Ok(Box::new(RunwayAdapter::new(config)?)), + "pika" => Ok(Box::new(PikaAdapter::new(config)?)), // 阶段 2 适配器占位 - "runway" | "pika" | "kling" | "edge-tts" | "elevenlabs" | "cartesia" | "deepgram" - | "assemblyai" => Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2 待实现)"), - }), + "kling" | "edge-tts" | "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { + Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 2 待实现)"), + }) + } // 未知 provider _ => Err(AibridgeError::provider_not_found(format!( "{provider}(未知 provider,支持:{})", @@ -422,6 +429,20 @@ mod tests { assert_eq!(adapter.provider_type(), "stability"); } + #[test] + fn create_runway_returns_adapter() { + // 阶段 2b:RunwayAdapter 自带 DEFAULT_RUNWAY_BASE_URL 兜底,仅需 api_key + let adapter = create_adapter(config_for("runway")).expect("工厂应能创建 runway 适配器"); + assert_eq!(adapter.provider_type(), "runway"); + } + + #[test] + fn create_pika_returns_adapter() { + // 阶段 2b:PikaAdapter 自带 DEFAULT_PIKA_BASE_URL 兜底,仅需 api_key + let adapter = create_adapter(config_for("pika")).expect("工厂应能创建 pika 适配器"); + assert_eq!(adapter.provider_type(), "pika"); + } + #[test] fn create_additional_models_aliases_map_to_main_provider_type() { // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: @@ -449,8 +470,8 @@ mod tests { #[test] fn create_phase2_adapter_returns_phase2_message() { - // runway 仍为阶段 2 占位(未实现),返 ProviderNotFound - let result = create_adapter(config_for("runway")); + // kling 仍为阶段 2 占位(未实现),返 ProviderNotFound + let result = create_adapter(config_for("kling")); if let Err(AibridgeError::ProviderNotFound { provider }) = result { assert!(provider.contains("阶段 2")); } else { @@ -492,6 +513,9 @@ mod tests { // 阶段 2b 第一批独立协议 assert!(is_known_provider("anthropic")); assert!(is_known_provider("stability")); + // 阶段 2b 第二批独立协议 + assert!(is_known_provider("runway")); + assert!(is_known_provider("pika")); } #[test] diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 5062c1b..0800bea 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -51,3 +51,9 @@ pub mod anthropic; /// Stability AI 适配器:阶段 2b 独立协议,文生图/图生图 pub mod stability; + +/// Runway 适配器:阶段 2b 独立协议,视频生成(文生视频/图生视频/任务轮询) +pub mod runway; + +/// Pika 适配器:阶段 2b 独立协议,视频生成(文生视频/图生视频/任务轮询) +pub mod pika; From ee89f3282fda5e01b56cc5cfff2015ac7aeb5566 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:12:06 +0800 Subject: [PATCH 39/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b=20kling=20=E9=80=82=E9=85=8D=E5=99=A8=EF=BC=88=E5=8F=AF?= =?UTF-8?q?=E7=81=B5=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 实现 KlingAdapter(独立协议,非 OpenAI 兼容): - video_create:POST /videos/generations(文生视频)/ /videos/image2video(图生视频) - video_poll:GET /videos/generations/{task_id} - list_models:硬编码 kling-v1 / v1-5 / v2 - 认证:Bearer Token;默认 base_url https://api.klingai.com/v1 - 状态映射:submitted/queued→pending、processing、succeed/success→success、failed/error→failed - 视频 URL:data.task_result.videos[0].url;错误:task_status_msg/error - 错误映射:401→Authentication、429→RateLimit、404→ModelNotFound、400→Validation、5xx→Api - 不支持能力(chat/image/embed/audio)走 trait 默认实现返 UnsupportedCapability - 71 个 mockito 单测覆盖正常/错误/边界路径 对应 Python v1 agn/adapters/kling.py。 --- crates/aibridge-core/src/adapters/kling.rs | 1749 ++++++++++++++++++++ 1 file changed, 1749 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/kling.rs diff --git a/crates/aibridge-core/src/adapters/kling.rs b/crates/aibridge-core/src/adapters/kling.rs new file mode 100644 index 0000000..eb7ef16 --- /dev/null +++ b/crates/aibridge-core/src/adapters/kling.rs @@ -0,0 +1,1749 @@ +//! Kling 适配器(可灵,视频生成,独立协议) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/kling.py`。 +//! +//! Kling(快手可灵)API 为**独立协议**(非 OpenAI 兼容),不复用 `OpenAiCompatAdapter`: +//! - 文生视频:`POST /videos/generations` +//! - 图生视频:`POST /videos/image2video` +//! - 查询任务状态:`GET /videos/generations/{task_id}` +//! - 认证:`Authorization: Bearer {api_key}` +//! - 默认 Base URL:`https://api.klingai.com/v1` +//! +//! ## 请求体规范 +//! - `model_name` / `prompt` 必选 +//! - 图生视频:`reference_images[0]` → `image` 字段 +//! - `negative_prompt` / `cfg_scale` / `duration` / `aspect_ratio` 可选透传 +//! - `camera_control`(镜头控制对象)/ `mode`(std/pro 档位)等厂商特有参数走 `extra` 透传 +//! +//! ## 响应解析 +//! - 创建/查询响应的任务信息包裹在 `data` 字段内(缺失时回退到顶层,更健壮) +//! - 任务 ID:`data.task_id`(缺失回退顶层,再缺失生成 `vid_` 前缀 ID) +//! - 状态:`data.task_status`,映射 submitted/queued → pending、processing、succeed/success → success、failed/error → failed +//! - 视频 URL:`data.task_result.videos[0].url` +//! - 错误信息:`data.task_status_msg` / `data.error` +//! - `progress`:success=100、pending=0、其余=50(与 Python v1 一致,Kling 响应无 progress 字段) +//! +//! ## 阶段范围 +//! 阶段 2b 实现 video_create + video_poll + list_models(硬编码)。 +//! 不支持的能力(chat / image / embed / audio)走 `Adapter` trait 默认实现返 +//! `UnsupportedCapability`,与 Python v1 抛 `UnsupportedCapabilityError` 行为一致。 +//! Python v1 明确不支持 image_generate(Kolors 是单独模型)。 + +use async_trait::async_trait; +use serde_json::{json, Value}; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::common::{ModelInfo, ModelType, TaskStatus, VideoMode}; +use crate::model::image::FileInput; +use crate::model::video::{VideoRequest, VideoStatus, VideoTask}; +use crate::util; + +/// Kling 默认 Base URL +/// +/// 对应 Python v1 `DEFAULT_BASE_URL`(已含 `/v1` 前缀)。 +pub const DEFAULT_KLING_BASE_URL: &str = "https://api.klingai.com/v1"; + +// ==================== 能力集合构造 ==================== + +/// Kling 支持的能力集合 +/// +/// 对齐 Python v1 `KlingAdapter.supported_capabilities = ["video"]`。 +/// Rust 用 `VideoGenerate` 表达视频生成能力(含 video_create + video_poll), +/// 并声明 text2video / image2video 两个子能力。 +fn kling_capabilities() -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::VideoGenerate); + caps.insert(Capabilities::VideoText2Video); + caps.insert(Capabilities::VideoImage2Video); + caps +} + +// ==================== Kling 视频生成适配器 ==================== + +/// Kling 适配器 +/// +/// 快手可灵视频生成平台,支持文生视频与图生视频(kling-v1 / v1-5 / v2)。 +/// 官方 API 文档:https://app.klingai.com/docs/api +/// +/// ## API 规范 +/// - Base URL: `https://api.klingai.com/v1` +/// - 文生视频: `POST /videos/generations` +/// - 图生视频: `POST /videos/image2video` +/// - 查询状态: `GET /videos/generations/{task_id}` +/// - 认证: `Authorization: Bearer {api_key}` +/// +/// ## 阶段范围 +/// 阶段 2b 实现 video_create + video_poll + list_models(硬编码)。 +pub struct KlingAdapter { + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// Provider 配置(api_key / base_url / timeout 等) + config: ProviderConfig, + /// 实际 base_url(已合并 config.base_url 与默认值) + base_url: String, + /// 支持的能力集合 + capabilities: CapabilitySet, +} + +impl KlingAdapter { + /// 创建 Kling 适配器 + /// + /// `config.base_url` 为空时用 [`DEFAULT_KLING_BASE_URL`] 兜底。 + /// `config.api_key` 为空时不在此处报错(由上层按 `requires_api_key` 校验)。 + pub fn new(config: ProviderConfig) -> Result { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_KLING_BASE_URL.to_string()); + + let opts = ClientOptions::builder() + .api_key(config.api_key.clone().unwrap_or_default()) + .base_url(base_url.clone()) + .timeout(config.timeout) + .max_retries(config.max_retries) + .retry_delay(config.retry_delay) + .build(); + let http = HttpClient::new(&opts)?; + + Ok(Self { + http, + config, + base_url, + capabilities: kling_capabilities(), + }) + } + + /// 用显式 HttpClient 构造(测试用,可注入 mockito 后端) + #[cfg(test)] + pub fn with_http(http: HttpClient, config: ProviderConfig) -> Self { + let base_url = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_KLING_BASE_URL.to_string()); + Self { + http, + config, + base_url, + capabilities: kling_capabilities(), + } + } + + /// API key(可能为空) + fn api_key(&self) -> &str { + self.config.api_key.as_deref().unwrap_or("") + } + + /// base_url + pub fn base_url(&self) -> &str { + &self.base_url + } + + /// 拼接完整 URL(base_url + 相对路径) + fn url(&self, path: &str) -> String { + let base = self.base_url.trim_end_matches('/'); + let path = path.trim_start_matches('/'); + format!("{base}/{path}") + } + + /// 校验请求的能力是否被支持(不支持则返 UnsupportedCapability) + fn ensure_capability(&self, cap: Capabilities) -> Result<()> { + if self.capabilities.contains(&cap) { + Ok(()) + } else { + Err(AibridgeError::UnsupportedCapability { + capability: format!("{} (provider: kling)", cap.as_str()), + }) + } + } + + /// 发送带 Bearer 认证的 POST JSON 请求,并用 Kling 错误映射处理响应 + async fn post_authed_json(&self, path: &str, body: &Value) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .post(&url) + .bearer_auth(self.api_key()) + .header("Content-Type", "application/json") + .header("Accept", "application/json") + .json(body) + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 发送带 Bearer 认证的 GET 请求,并用 Kling 错误映射处理响应 + async fn get_authed_json(&self, path: &str) -> Result { + let url = self.url(path); + let resp = self + .http + .inner() + .get(&url) + .bearer_auth(self.api_key()) + .header("Accept", "application/json") + .send() + .await + .map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body_text = resp.text().await.unwrap_or_default(); + return Err(Self::map_api_error(status_code, &body_text)); + } + resp.json::().await.map_err(AibridgeError::from) + } + + /// 构造 Kling 视频生成请求体 + /// + /// 移植自 Python v1 `video_create`: + /// - `model_name` / `prompt` 必选(Kling 用 `model_name` 而非 `model`) + /// - image2video 模式且 `reference_images[0]` 可转 URL 时,写入 `image` 字段 + /// - `negative_prompt` / `cfg_scale` / `duration` / `aspect_ratio` 可选透传 + /// - `camera_control`(镜头控制)/ `mode`(std/pro)等厂商特有参数走 `extra` 透传到顶层 + fn build_video_body(&self, req: &VideoRequest) -> Value { + let mut body = json!({ + "model_name": req.model, + "prompt": req.prompt, + }); + + // image2video 模式:reference_images[0] → image 字段 + if is_image2video_request(req) { + if let Some(url) = req.reference_images.first().and_then(file_input_to_url) { + body["image"] = json!(url); + } + } + + if let Some(np) = &req.negative_prompt { + body["negative_prompt"] = json!(np); + } + if let Some(cfg) = req.cfg_scale { + body["cfg_scale"] = json!(cfg); + } + if let Some(d) = req.duration { + body["duration"] = json!(d); + } + if let Some(ar) = &req.aspect_ratio { + body["aspect_ratio"] = json!(ar); + } + + // extra 透传(合并到顶层:camera_control / mode 等厂商特有参数) + if let Some(obj) = body.as_object_mut() { + for (k, v) in &req.extra { + obj.insert(k.clone(), v.clone()); + } + } + body + } + + /// 解析 Kling 创建响应 → VideoTask + /// + /// 移植自 Python v1 `video_create` 响应解析: + /// - 任务信息包裹在 `data` 字段内(缺失时回退顶层,更健壮) + /// - 任务 ID:`data.task_id`(回退顶层 task_id,再缺失生成 `vid_` 前缀 ID) + /// - 状态:`data.task_status`(缺失默认 pending) + /// - `created_at`:优先取响应值,缺失用当前时间戳(对齐 Python v1) + fn parse_video_task(value: &Value, model: &str) -> Result { + let data = value.get("data").unwrap_or(value); + let task_id = data + .get("task_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + value + .get("task_id") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }) + .unwrap_or_else(|| util::generate_id("vid")); + let raw_status = data + .get("task_status") + .and_then(|v| v.as_str()) + .or_else(|| value.get("task_status").and_then(|v| v.as_str())) + .unwrap_or("pending"); + let created_at = data + .get("created_at") + .and_then(|v| v.as_u64()) + .or_else(|| value.get("created_at").and_then(|v| v.as_u64())) + .unwrap_or_else(util::current_timestamp); + Ok(VideoTask { + task_id, + model: model.to_string(), + status: map_kling_status(raw_status), + created_at, + }) + } + + /// 解析 Kling 查询响应 → VideoStatus + /// + /// 移植自 Python v1 `video_poll`: + /// - 任务信息包裹在 `data` 字段内(缺失时回退顶层) + /// - 视频 URL:`data.task_result.videos[0].url` + /// - 错误信息:`data.task_status_msg` / `data.error` + /// - `progress`:success=100、pending=0、其余=50(Kling 响应无 progress 字段,按状态推算) + /// - `updated_at` 缺失时回退当前时间戳(对齐 Python v1) + fn parse_video_status(value: &Value, task_id: &str) -> VideoStatus { + let data = value.get("data").unwrap_or(value); + let raw_status = data + .get("task_status") + .and_then(|v| v.as_str()) + .unwrap_or(""); + let status = map_kling_status(raw_status); + + // 视频 URL:data.task_result.videos[0].url + let video_url = data + .get("task_result") + .and_then(|tr| tr.get("videos")) + .and_then(|v| v.as_array()) + .and_then(|arr| arr.first()) + .and_then(|f| f.get("url")) + .and_then(|v| v.as_str()) + .map(str::to_owned); + + // 错误信息:task_status_msg 优先,回退 error + let error = data + .get("task_status_msg") + .and_then(|v| v.as_str()) + .map(str::to_owned) + .or_else(|| { + data.get("error") + .and_then(|v| v.as_str()) + .map(str::to_owned) + }); + + // progress:按状态推算(Kling 响应无 progress 字段) + let progress = match status { + TaskStatus::Success => Some(100), + TaskStatus::Pending => Some(0), + _ => Some(50), + }; + + let created_at = data.get("created_at").and_then(|v| v.as_u64()); + let updated_at = data + .get("updated_at") + .and_then(|v| v.as_u64()) + .or_else(|| Some(util::current_timestamp())); + + VideoStatus { + task_id: task_id.to_string(), + status, + video_url, + progress, + error, + created_at, + updated_at, + } + } + + /// 将 Kling API 错误响应映射为 AibridgeError + /// + /// 移植自 Python v1 `_handle_kling_error` 并按阶段 2b 统一错误映射要求调整: + /// - 401 → Authentication("Invalid Kling API key") + /// - 429 → RateLimit("Kling rate limit exceeded or quota exhausted") + /// - 404 → ModelNotFound(Kling 任务/资源不存在) + /// - 400 → Validation(请求参数校验错误) + /// - 其余 ≥400 → Api(提取 error.message / message / error / detail,回退 `HTTP {status}`) + pub fn map_api_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::Authentication { + message: "Invalid Kling API key".to_string(), + }, + 429 => AibridgeError::RateLimit { + message: "Kling rate limit exceeded or quota exhausted".to_string(), + retry_after: None, + }, + 404 => AibridgeError::ModelNotFound { + model: "Kling task or resource not found".to_string(), + }, + 400 => AibridgeError::validation_with_details( + parse_error_message(body, status), + serde_json::from_str(body).unwrap_or(serde_json::Value::Null), + ), + _ => { + let message = parse_error_message(body, status); + AibridgeError::Api { status, message } + } + } + } +} + +#[async_trait] +impl Adapter for KlingAdapter { + fn provider_type(&self) -> &str { + "kling" + } + + fn provider_name(&self) -> &str { + "可灵 Kling" + } + + fn capabilities(&self) -> CapabilitySet { + self.capabilities.clone() + } + + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HttpClient 在 new() 时已构造,无额外资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // HttpClient 由 Drop 自动释放,无额外资源 + Ok(()) + } + + /// 创建视频生成任务 + /// + /// - image2video 模式且有可转 URL 的 reference_images[0] → `POST /videos/image2video` + /// - 其余 → `POST /videos/generations` + async fn video_create(&self, req: VideoRequest) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let body = self.build_video_body(&req); + let endpoint = if is_image2video_request(&req) { + "videos/image2video" + } else { + "videos/generations" + }; + let value = self.post_authed_json(endpoint, &body).await?; + Self::parse_video_task(&value, &req.model) + } + + /// 查询视频任务状态:`GET /videos/generations/{task_id}` + async fn video_poll(&self, task_id: &str, _model: &str) -> Result { + self.ensure_capability(Capabilities::VideoGenerate)?; + let path = format!("videos/generations/{task_id}"); + let value = self.get_authed_json(&path).await?; + Ok(Self::parse_video_status(&value, task_id)) + } + + /// 模型列表(硬编码) + /// + /// Kling 无标准 `/models` 端点,暂保留硬编码列表(与 Python v1 一致)。 + /// 含 kling-v1 / kling-v1-5 / kling-v2 三个模型。 + async fn list_models(&self, filter: Option) -> Result> { + let models = kling_hardcoded_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } + + // chat / chat_stream / image / embed / audio 走 trait 默认实现返 UnsupportedCapability, + // 与 Python v1 抛 UnsupportedCapabilityError 行为一致。 +} + +// ==================== 内部:辅助函数 ==================== + +/// 判断是否为图生视频请求 +/// +/// 同时满足:mode 为 Image2Video,且 reference_images[0] 可转 URL(Url/Base64)。 +/// 对齐 Python v1 `mode == "image2video" and reference_images` 的端点选择逻辑。 +fn is_image2video_request(req: &VideoRequest) -> bool { + matches!(req.mode, VideoMode::Image2Video) + && req + .reference_images + .first() + .and_then(file_input_to_url) + .is_some() +} + +/// 映射 Kling 状态字符串到统一 TaskStatus +/// +/// 移植自 Python v1 `_map_kling_status`(大小写不敏感,比 Python 更健壮): +/// - submitted / queued → Pending +/// - processing → Processing +/// - succeed / success → Success +/// - failed / error → Failed +/// - 未知 → Pending(与 Python 默认值一致) +fn map_kling_status(raw: &str) -> TaskStatus { + match raw.to_lowercase().as_str() { + "submitted" | "queued" => TaskStatus::Pending, + "processing" => TaskStatus::Processing, + "succeed" | "success" => TaskStatus::Success, + "failed" | "error" => TaskStatus::Failed, + _ => TaskStatus::Pending, + } +} + +/// 从 FileInput 提取 URL 字符串 +/// +/// Kling 的 `image` 字段接受 URL 或 base64 字符串: +/// - `Url(s)` / `Base64(s)` → `Some(s)` +/// - `Path(_)` / `Bytes(_)` → `None`(需调用方先上传为可访问 URL) +fn file_input_to_url(input: &FileInput) -> Option { + match input { + FileInput::Url(s) | FileInput::Base64(s) => Some(s.clone()), + FileInput::Path(_) | FileInput::Bytes(_) => None, + } +} + +/// 解析错误体中的 message 字段 +/// +/// 通用错误消息提取,兼容多种错误体结构: +/// - `{"error": {"message": "..."}}`(OpenAI 风格) +/// - `{"error": "..."}`(顶层 error 字符串) +/// - `{"message": "..."}`(顶层 message,Kling 常用) +/// - `{"detail": "..."}`(顶层 detail) +/// +/// 解析失败时回退到 `HTTP {status}` 字符串。 +fn parse_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + // error.message(OpenAI 风格) + if let Some(msg) = v + .get("error") + .and_then(|e| e.get("message")) + .and_then(|m| m.as_str()) + { + return msg.to_string(); + } + // 顶层 error 字符串 + if let Some(msg) = v.get("error").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 message(Kling 常用) + if let Some(msg) = v.get("message").and_then(|m| m.as_str()) { + return msg.to_string(); + } + // 顶层 detail + if let Some(msg) = v.get("detail").and_then(|m| m.as_str()) { + return msg.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + format!("HTTP {status}: {body}") + } +} + +/// Kling 硬编码模型列表 +/// +/// 对应 Python v1 `KlingAdapter.list_models`。 +/// 注意:该 Provider 无标准 `/models` 端点,暂保留硬编码列表。 +fn kling_hardcoded_models() -> Vec { + vec![ + ModelInfo { + id: "kling-v1".into(), + name: "Kling 1.0".into(), + model_type: ModelType::Video, + provider: "kling".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Kling 1.0 标准版本".into()), + created: None, + }, + ModelInfo { + id: "kling-v1-5".into(), + name: "Kling 1.5".into(), + model_type: ModelType::Video, + provider: "kling".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Kling 1.5 改进版本".into()), + created: None, + }, + ModelInfo { + id: "kling-v2".into(), + name: "Kling 2.0".into(), + model_type: ModelType::Video, + provider: "kling".into(), + capabilities: vec!["text2video".into(), "image2video".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Kling 2.0 最新版本,质量更好".into()), + created: None, + }, + ] +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::chat::ChatRequest; + use crate::model::image::ImageRequest; + use mockito::Server; + use serde_json::json; + + // ==================== 通用测试辅助 ==================== + + /// 构造测试用 KlingAdapter(指向 mockito server) + fn make_kling(server: &Server) -> KlingAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(server.url()) + .timeout(5) + .build(); + let config = ProviderConfig::from_options("kling", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url(server.url()).build()).unwrap(); + KlingAdapter::with_http(http, config) + } + + /// 构造不指向任何 server 的 KlingAdapter(用于不发请求的元信息/能力测试) + fn make_kling_no_server() -> KlingAdapter { + let opts = ClientOptions::builder() + .api_key("test-key") + .base_url(DEFAULT_KLING_BASE_URL) + .build(); + let config = ProviderConfig::from_options("kling", opts); + KlingAdapter::new(config).expect("KlingAdapter 构造应成功") + } + + // ============ 元信息 ============ + + #[test] + fn kling_provider_type_and_name_match_python() { + let adapter = make_kling_no_server(); + assert_eq!(adapter.provider_type(), "kling"); + assert_eq!(adapter.provider_name(), "可灵 Kling"); + } + + #[test] + fn kling_requires_api_key_is_true() { + let adapter = make_kling_no_server(); + assert!(adapter.requires_api_key()); + } + + #[test] + fn kling_capabilities_contains_only_video() { + let adapter = make_kling_no_server(); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::VideoGenerate)); + assert!(caps.contains(&Capabilities::VideoText2Video)); + assert!(caps.contains(&Capabilities::VideoImage2Video)); + // chat / image 不声明 + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn kling_base_url_defaults_when_missing() { + let opts = ClientOptions::builder().api_key("k").build(); + let config = ProviderConfig::from_options("kling", opts); + let adapter = KlingAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_KLING_BASE_URL); + } + + #[test] + fn kling_base_url_uses_config_when_provided() { + let opts = ClientOptions::builder() + .api_key("k") + .base_url("https://custom.kling-proxy.com/v1") + .build(); + let config = ProviderConfig::from_options("kling", opts); + let adapter = KlingAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), "https://custom.kling-proxy.com/v1"); + } + + #[test] + fn kling_base_url_ignores_empty_string() { + // 空白 base_url 应回退到默认值 + let opts = ClientOptions::builder() + .api_key("k") + .base_url(" ") + .build(); + let config = ProviderConfig::from_options("kling", opts); + let adapter = KlingAdapter::new(config).unwrap(); + assert_eq!(adapter.base_url(), DEFAULT_KLING_BASE_URL); + } + + // ============ video_create 正常路径 ============ + + #[tokio::test] + async fn kling_video_create_success_returns_task() { + let mut server = Server::new_async().await; + let body = json!({ + "code": 0, + "message": "success", + "data": { + "task_id": "vid-abc123", + "task_status": "submitted" + } + }); + let mock = server + .mock("POST", "/videos/generations") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "A cat running in the park").build(); + let task = adapter + .video_create(req) + .await + .expect("video_create 应成功"); + + assert_eq!(task.task_id, "vid-abc123"); + assert_eq!(task.model, "kling-v1"); + // submitted → pending + assert_eq!(task.status, TaskStatus::Pending); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_sends_model_name_and_prompt() { + // 验证请求体用 model_name 字段(非 model),且包含 prompt + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "model_name": "kling-v1", + "prompt": "A cat running" + }))) + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "x", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "A cat running").build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_passes_optional_params() { + // negative_prompt / cfg_scale / duration / aspect_ratio 透传 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "negative_prompt": "blurry, low quality", + "cfg_scale": 0.5, + "duration": 10, + "aspect_ratio": "16:9" + }))) + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "x", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v2", "a beautiful sunset") + .negative_prompt("blurry, low quality") + .cfg_scale(0.5) + .duration(10) + .aspect_ratio("16:9") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_image2video_uses_image2video_endpoint() { + // image2video 模式:走 /videos/image2video 端点,body 含 image 字段 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/image2video") + .match_body(mockito::Matcher::PartialJson(json!({ + "image": "https://example.com/input.jpg" + }))) + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "vid-img2vid-001", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1-5", "Make this image move") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/input.jpg")]) + .build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "vid-img2vid-001"); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_text2video_does_not_send_image_field() { + // text2video 模式即使有 reference_images 也不发送 image 字段,且走 /videos/generations + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/generations") + .match_body(mockito::Matcher::Json(json!({ + "model_name": "kling-v1", + "prompt": "a cat" + }))) + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "x", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat") + .mode(VideoMode::Text2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_image2video_without_reference_falls_back_to_generations() { + // image2video 模式但无 reference_images → 回退到 /videos/generations 端点 + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/generations") + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "x", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat") + .mode(VideoMode::Image2Video) + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_passes_extra_params() { + // extra 字段透传到顶层(camera_control / mode 等厂商特有参数) + let mut server = Server::new_async().await; + let mock = server + .mock("POST", "/videos/generations") + .match_body(mockito::Matcher::PartialJson(json!({ + "camera_control": {"pan": "left"}, + "mode": "std" + }))) + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "x", "task_status": "submitted"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat") + .extra("camera_control", json!({"pan": "left"})) + .extra("mode", "std") + .build(); + let _ = adapter.video_create(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_create_accepts_top_level_task_id() { + // 任务 ID 在顶层 task_id(无 data 包裹) + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(200) + .with_body(json!({"task_id": "top-level-id", "task_status": "queued"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.task_id, "top-level-id"); + // queued → pending + assert_eq!(task.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn kling_video_create_uses_generated_id_when_missing() { + // 响应缺任务 ID 时,回退到生成的 vid_ ID + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(200) + .with_body(json!({"code": 0, "data": {"task_status": "submitted"}}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert!(task.task_id.starts_with("vid_")); + } + + #[tokio::test] + async fn kling_video_create_processing_status() { + // processing 状态映射 + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "p-1", "task_status": "processing"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let task = adapter.video_create(req).await.unwrap(); + assert_eq!(task.status, TaskStatus::Processing); + } + + // ============ video_create 错误路径 ============ + + #[tokio::test] + async fn kling_video_create_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(401) + .with_body(json!({"message": "invalid key"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Authentication { message } => assert!(message.contains("Kling")), + _ => panic!("应为 Authentication"), + } + } + + #[tokio::test] + async fn kling_video_create_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(429) + .with_body(json!({"message": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::RateLimit { + message, + retry_after, + } => { + assert!(message.contains("Kling")); + assert!(retry_after.is_none()); + } + _ => panic!("应为 RateLimit"), + } + } + + #[tokio::test] + async fn kling_video_create_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(404) + .with_body(json!({"message": "not found"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn kling_video_create_error_400_returns_validation() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(400) + .with_body(json!({"error": {"message": "invalid prompt"}}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Validation { message, .. } => { + assert!(message.contains("invalid prompt")); + } + _ => panic!("应为 Validation"), + } + } + + #[tokio::test] + async fn kling_video_create_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(500) + .with_body(json!({"error": {"message": "internal"}}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert!(message.contains("internal")); + } + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn kling_video_create_error_no_json_body_falls_back() { + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(502) + .with_body("Bad Gateway") + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + #[tokio::test] + async fn kling_video_create_error_extracts_top_level_message() { + // Kling 错误体常用顶层 message 字段 + let mut server = Server::new_async().await; + server + .mock("POST", "/videos/generations") + .with_status(503) + .with_body(json!({"message": "service unavailable"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let err = adapter.video_create(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "service unavailable"), + _ => panic!("应为 Api"), + } + } + + // ============ video_poll 正常路径 ============ + + #[tokio::test] + async fn kling_video_poll_success_returns_video_url() { + let mut server = Server::new_async().await; + let body = json!({ + "code": 0, + "data": { + "task_id": "vid-abc123", + "task_status": "succeed", + "task_result": { + "videos": [{"url": "https://cdn.example.com/video.mp4"}] + }, + "created_at": 1700000000, + "updated_at": 1700000300 + } + }); + let mock = server + .mock("GET", "/videos/generations/vid-abc123") + .match_header("authorization", "Bearer test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter + .video_poll("vid-abc123", "kling-v1") + .await + .expect("video_poll 应成功"); + + assert_eq!(status.task_id, "vid-abc123"); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://cdn.example.com/video.mp4") + ); + assert_eq!(status.progress, Some(100)); + assert_eq!(status.created_at, Some(1700000000)); + assert_eq!(status.updated_at, Some(1700000300)); + assert!(status.error.is_none()); + mock.assert_async().await; + } + + #[tokio::test] + async fn kling_video_poll_pending_status() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-1") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-1", + "task_status": "submitted", + "created_at": 1700000000, + "updated_at": 1700000100 + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-1", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + assert_eq!(status.progress, Some(0)); + assert!(status.video_url.is_none()); + } + + #[tokio::test] + async fn kling_video_poll_processing_status() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-2") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-2", + "task_status": "processing", + "created_at": 1700000000, + "updated_at": 1700000200 + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-2", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Processing); + assert_eq!(status.progress, Some(50)); + } + + #[tokio::test] + async fn kling_video_poll_failed_returns_error_msg() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-3") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-3", + "task_status": "failed", + "task_status_msg": "Insufficient credits", + "created_at": 1700000000, + "updated_at": 1700000300 + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-3", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("Insufficient credits")); + // failed → progress 50 + assert_eq!(status.progress, Some(50)); + } + + #[tokio::test] + async fn kling_video_poll_failed_extracts_error_field() { + // task_status_msg 缺失时回退 error 字段 + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-4") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-4", + "task_status": "error", + "error": "internal failure" + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-4", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Failed); + assert_eq!(status.error.as_deref(), Some("internal failure")); + } + + #[tokio::test] + async fn kling_video_poll_queued_maps_to_pending() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-5") + .with_status(200) + .with_body( + json!({"code": 0, "data": {"task_id": "vid-5", "task_status": "queued"}}) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-5", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Pending); + } + + #[tokio::test] + async fn kling_video_poll_success_status_keyword() { + // "success" 关键字也映射到 Success + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-6") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-6", + "task_status": "success", + "task_result": {"videos": [{"url": "https://example.com/v.mp4"}]} + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-6", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://example.com/v.mp4") + ); + } + + #[tokio::test] + async fn kling_video_poll_accepts_top_level_data() { + // 无 data 包裹时,回退到顶层解析 + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-7") + .with_status(200) + .with_body( + json!({ + "task_id": "vid-7", + "task_status": "succeed", + "task_result": {"videos": [{"url": "https://top.example.com/v.mp4"}]} + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-7", "kling-v1").await.unwrap(); + assert_eq!(status.status, TaskStatus::Success); + assert_eq!( + status.video_url.as_deref(), + Some("https://top.example.com/v.mp4") + ); + } + + #[tokio::test] + async fn kling_video_poll_updated_at_falls_back_to_current_timestamp() { + // updated_at 缺失时回退当前时间戳(对齐 Python v1) + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-8") + .with_status(200) + .with_body( + json!({ + "code": 0, + "data": { + "task_id": "vid-8", + "task_status": "processing", + "created_at": 1700000000 + } + }) + .to_string(), + ) + .create_async() + .await; + + let adapter = make_kling(&server); + let status = adapter.video_poll("vid-8", "kling-v1").await.unwrap(); + assert_eq!(status.created_at, Some(1700000000)); + assert!(status.updated_at.is_some()); + } + + // ============ video_poll 错误路径 ============ + + #[tokio::test] + async fn kling_video_poll_error_401_returns_authentication() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-x") + .with_status(401) + .with_body(json!({"message": "bad key"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let err = adapter.video_poll("vid-x", "kling-v1").await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn kling_video_poll_error_404_returns_model_not_found() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/nonexistent") + .with_status(404) + .with_body(json!({"message": "task not found"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let err = adapter + .video_poll("nonexistent", "kling-v1") + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::ModelNotFound { .. })); + } + + #[tokio::test] + async fn kling_video_poll_error_429_returns_rate_limit() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-y") + .with_status(429) + .with_body(json!({"message": "slow down"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let err = adapter.video_poll("vid-y", "kling-v1").await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn kling_video_poll_error_500_returns_api() { + let mut server = Server::new_async().await; + server + .mock("GET", "/videos/generations/vid-z") + .with_status(500) + .with_body(json!({"message": "internal"}).to_string()) + .create_async() + .await; + + let adapter = make_kling(&server); + let err = adapter.video_poll("vid-z", "kling-v1").await.unwrap_err(); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert!(message.contains("internal")); + } + _ => panic!("应为 Api"), + } + } + + // ============ 不支持的能力 ============ + + #[tokio::test] + async fn kling_chat_returns_unsupported() { + let adapter = make_kling_no_server(); + let req = ChatRequest::builder("kling-v1", vec![]).build(); + let err = adapter.chat(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn kling_image_generate_returns_unsupported() { + let adapter = make_kling_no_server(); + let req = ImageRequest::builder("kling-v1", "a cat").build(); + let err = adapter.image_generate(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + #[tokio::test] + async fn kling_embed_returns_unsupported() { + let adapter = make_kling_no_server(); + let req = crate::model::options::EmbedRequest { + model: "kling-v1".into(), + input: crate::model::options::EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: std::collections::HashMap::new(), + }; + let err = adapter.embed(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::UnsupportedCapability { .. })); + } + + // ============ list_models ============ + + #[tokio::test] + async fn kling_list_models_returns_hardcoded() { + let adapter = make_kling_no_server(); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + assert_eq!(models[0].id, "kling-v1"); + assert_eq!(models[0].provider, "kling"); + assert_eq!(models[0].model_type, ModelType::Video); + assert_eq!(models[1].id, "kling-v1-5"); + assert_eq!(models[2].id, "kling-v2"); + } + + #[tokio::test] + async fn kling_list_models_filter_by_video_type() { + let adapter = make_kling_no_server(); + let videos = adapter.list_models(Some(ModelType::Video)).await.unwrap(); + assert_eq!(videos.len(), 3); + assert!(videos.iter().all(|m| m.model_type == ModelType::Video)); + } + + #[tokio::test] + async fn kling_list_models_filter_by_image_returns_empty() { + let adapter = make_kling_no_server(); + let images = adapter.list_models(Some(ModelType::Image)).await.unwrap(); + assert!(images.is_empty()); + } + + #[tokio::test] + async fn kling_list_models_filter_by_chat_returns_empty() { + let adapter = make_kling_no_server(); + let chats = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chats.is_empty()); + } + + // ============ start / close ============ + + #[tokio::test] + async fn kling_start_and_close_are_noops() { + let mut adapter = make_kling_no_server(); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 错误映射单元测试 ============ + + #[test] + fn map_api_error_401_is_authentication() { + let err = KlingAdapter::map_api_error(401, "{\"message\":\"bad\"}"); + match err { + AibridgeError::Authentication { message } => assert!(message.contains("Kling")), + _ => panic!("应为 Authentication"), + } + } + + #[test] + fn map_api_error_429_is_rate_limit() { + let err = KlingAdapter::map_api_error(429, "{}"); + match err { + AibridgeError::RateLimit { + message, + retry_after, + } => { + assert!(message.contains("Kling")); + assert!(retry_after.is_none()); + } + _ => panic!("应为 RateLimit"), + } + } + + #[test] + fn map_api_error_404_is_model_not_found() { + let err = KlingAdapter::map_api_error(404, "{}"); + match err { + AibridgeError::ModelNotFound { model } => assert!(model.contains("Kling")), + _ => panic!("应为 ModelNotFound"), + } + } + + #[test] + fn map_api_error_400_is_validation() { + let err = KlingAdapter::map_api_error(400, "{\"error\":{\"message\":\"bad param\"}}"); + match err { + AibridgeError::Validation { message, .. } => assert_eq!(message, "bad param"), + _ => panic!("应为 Validation"), + } + } + + #[test] + fn map_api_error_500_extracts_message() { + let err = KlingAdapter::map_api_error(500, "{\"error\":{\"message\":\"internal\"}}"); + match err { + AibridgeError::Api { status, message } => { + assert_eq!(status, 500); + assert_eq!(message, "internal"); + } + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_extracts_top_level_message() { + // Kling 错误体常用顶层 message + let err = KlingAdapter::map_api_error(503, "{\"message\":\"unavailable\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "unavailable"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_extracts_detail_field() { + let err = KlingAdapter::map_api_error(422, "{\"detail\":\"validation failed\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "validation failed"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_extracts_top_level_error_string() { + let err = KlingAdapter::map_api_error(409, "{\"error\":\"conflict\"}"); + match err { + AibridgeError::Api { message, .. } => assert_eq!(message, "conflict"), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_no_json_falls_back_to_http_status() { + let err = KlingAdapter::map_api_error(502, "Bad Gateway"); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("502")), + _ => panic!("应为 Api"), + } + } + + #[test] + fn map_api_error_empty_body_falls_back_to_http_status() { + let err = KlingAdapter::map_api_error(500, ""); + match err { + AibridgeError::Api { message, .. } => assert!(message.contains("500")), + _ => panic!("应为 Api"), + } + } + + // ============ map_kling_status 单元测试 ============ + + #[test] + fn map_kling_status_pending_variants() { + assert_eq!(map_kling_status("submitted"), TaskStatus::Pending); + assert_eq!(map_kling_status("queued"), TaskStatus::Pending); + } + + #[test] + fn map_kling_status_processing() { + assert_eq!(map_kling_status("processing"), TaskStatus::Processing); + } + + #[test] + fn map_kling_status_success_variants() { + assert_eq!(map_kling_status("succeed"), TaskStatus::Success); + assert_eq!(map_kling_status("success"), TaskStatus::Success); + } + + #[test] + fn map_kling_status_failed_variants() { + assert_eq!(map_kling_status("failed"), TaskStatus::Failed); + assert_eq!(map_kling_status("error"), TaskStatus::Failed); + } + + #[test] + fn map_kling_status_case_insensitive() { + assert_eq!(map_kling_status("SUBMITTED"), TaskStatus::Pending); + assert_eq!(map_kling_status("Processing"), TaskStatus::Processing); + assert_eq!(map_kling_status("SUCCEED"), TaskStatus::Success); + assert_eq!(map_kling_status("FAILED"), TaskStatus::Failed); + } + + #[test] + fn map_kling_status_unknown_defaults_to_pending() { + assert_eq!(map_kling_status("unknown_state"), TaskStatus::Pending); + assert_eq!(map_kling_status(""), TaskStatus::Pending); + } + + // ============ file_input_to_url 单元测试 ============ + + #[test] + fn file_input_to_url_returns_url_for_url_variant() { + let f = FileInput::url("https://example.com/x.png"); + assert_eq!( + file_input_to_url(&f), + Some("https://example.com/x.png".to_string()) + ); + } + + #[test] + fn file_input_to_url_returns_base64_for_base64_variant() { + let f = FileInput::base64("aGVsbG8="); + assert_eq!(file_input_to_url(&f), Some("aGVsbG8=".to_string())); + } + + #[test] + fn file_input_to_url_returns_none_for_path_and_bytes() { + assert_eq!(file_input_to_url(&FileInput::path("/tmp/x")), None); + assert_eq!(file_input_to_url(&FileInput::bytes(vec![1, 2])), None); + } + + // ============ is_image2video_request 单元测试 ============ + + #[test] + fn is_image2video_request_true_for_image2video_with_url() { + let req = VideoRequest::builder("kling-v1", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + assert!(is_image2video_request(&req)); + } + + #[test] + fn is_image2video_request_false_for_text2video() { + let req = VideoRequest::builder("kling-v1", "a cat") + .mode(VideoMode::Text2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + assert!(!is_image2video_request(&req)); + } + + #[test] + fn is_image2video_request_false_for_image2video_without_reference() { + let req = VideoRequest::builder("kling-v1", "a cat") + .mode(VideoMode::Image2Video) + .build(); + assert!(!is_image2video_request(&req)); + } + + #[test] + fn is_image2video_request_false_for_image2video_with_path_input() { + // Path 类型无法转 URL,回退到 text2video 端点 + let req = VideoRequest::builder("kling-v1", "a cat") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::path("/tmp/x.png")]) + .build(); + assert!(!is_image2video_request(&req)); + } + + // ============ build_video_body 单元测试 ============ + + #[test] + fn build_video_body_includes_model_name_and_prompt() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("kling", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = KlingAdapter::with_http(http, config); + let req = VideoRequest::builder("kling-v1", "a cat").build(); + let body = adapter.build_video_body(&req); + assert_eq!(body["model_name"], "kling-v1"); + assert_eq!(body["prompt"], "a cat"); + // text2video 不发 image 字段 + assert!(body.get("image").is_none()); + } + + #[test] + fn build_video_body_image2video_sets_image_field() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("kling", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = KlingAdapter::with_http(http, config); + let req = VideoRequest::builder("kling-v1", "animate") + .mode(VideoMode::Image2Video) + .reference_images(vec![FileInput::url("https://example.com/a.png")]) + .build(); + let body = adapter.build_video_body(&req); + assert_eq!(body["image"], "https://example.com/a.png"); + } + + #[test] + fn build_video_body_passes_optional_fields() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("kling", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = KlingAdapter::with_http(http, config); + let req = VideoRequest::builder("kling-v2", "a cat") + .negative_prompt("blurry") + .cfg_scale(0.5) + .duration(10) + .aspect_ratio("16:9") + .build(); + let body = adapter.build_video_body(&req); + assert_eq!(body["negative_prompt"], "blurry"); + assert_eq!(body["cfg_scale"], 0.5); + assert_eq!(body["duration"], 10); + assert_eq!(body["aspect_ratio"], "16:9"); + } + + #[test] + fn build_video_body_extra_passthrough() { + let opts = ClientOptions::builder().base_url("https://x").build(); + let config = ProviderConfig::from_options("kling", opts); + let http = + HttpClient::new(&ClientOptions::builder().base_url("https://x").build()).unwrap(); + let adapter = KlingAdapter::with_http(http, config); + let req = VideoRequest::builder("kling-v1", "a cat") + .extra("camera_control", json!({"pan": "left"})) + .extra("mode", "std") + .build(); + let body = adapter.build_video_body(&req); + assert_eq!(body["camera_control"]["pan"], "left"); + assert_eq!(body["mode"], "std"); + } +} From d6493c7d0ad389b5ef4d923d2591136070ab809a Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:23:50 +0800 Subject: [PATCH 40/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20edge-tts=20=E9=80=82=E9=85=8D=E5=99=A8=EF=BC=88=E5=85=8D?= =?UTF-8?q?=E8=B4=B9=20TTS=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- Cargo.lock | 217 ++- crates/aibridge-core/Cargo.toml | 5 + crates/aibridge-core/src/adapters/edge_tts.rs | 1669 +++++++++++++++++ 3 files changed, 1880 insertions(+), 11 deletions(-) create mode 100644 crates/aibridge-core/src/adapters/edge_tts.rs diff --git a/Cargo.lock b/Cargo.lock index dfc9e0b..33f5b9a 100644 --- a/Cargo.lock +++ b/Cargo.lock @@ -20,14 +20,17 @@ dependencies = [ "base64", "bytes", "futures", + "hmac", "mockito", "once_cell", "rand 0.8.6", "reqwest", "serde", "serde_json", + "sha2", "thiserror 1.0.69", "tokio", + "tokio-tungstenite", "tracing", ] @@ -180,12 +183,27 @@ version = "2.13.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "b4388bee8683e3d04af747c73422af53102d2bd24d9eadb6cbc100baef4b43f8" +[[package]] +name = "block-buffer" +version = "0.10.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "3078c7629b62d3f0439517fa394996acacc5cbc91c5a20d8c658e77abd503a71" +dependencies = [ + "generic-array", +] + [[package]] name = "bumpalo" version = "3.20.3" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "72f5acc6cb2ba439de613abc23857ec3d78374d8ed5ac84e9d11336e87da8649" +[[package]] +name = "byteorder" +version = "1.5.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "1fd0f2584146f6f2ef48085050886acf353beff7305ebd1ae69500e27c67f64b" + [[package]] name = "bytes" version = "1.12.0" @@ -240,7 +258,7 @@ source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "d524456ba66e72eb8b115ff89e01e497f8e6d11d78b70b1aa13c0fbd97540a81" dependencies = [ "cfg-if", - "cpufeatures", + "cpufeatures 0.3.0", "rand_core 0.10.1", ] @@ -295,6 +313,15 @@ dependencies = [ "unicode-segmentation", ] +[[package]] +name = "cpufeatures" +version = "0.2.17" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "59ed5838eebb26a2bb2e58f6d5b5316989ae9d08bab10e0e6d103e656d1b0280" +dependencies = [ + "libc", +] + [[package]] name = "cpufeatures" version = "0.3.0" @@ -304,6 +331,16 @@ dependencies = [ "libc", ] +[[package]] +name = "crypto-common" +version = "0.1.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "78c8292055d1c1df0cce5d180393dc8cce0abec0a7102adb6c7b1eef6016d60a" +dependencies = [ + "generic-array", + "typenum", +] + [[package]] name = "ctor" version = "0.2.9" @@ -314,6 +351,23 @@ dependencies = [ "syn", ] +[[package]] +name = "data-encoding" +version = "2.11.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a4ae5f15dda3c708c0ade84bfee31ccab44a3da4f88015ed22f63732abe300c8" + +[[package]] +name = "digest" +version = "0.10.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ed9a281f7bc9b7576e61468ba615a66a5c8cfdff42420a70aa82701a3b1e292" +dependencies = [ + "block-buffer", + "crypto-common", + "subtle", +] + [[package]] name = "displaydoc" version = "0.2.6" @@ -456,6 +510,16 @@ dependencies = [ "slab", ] +[[package]] +name = "generic-array" +version = "0.14.7" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "85649ca51fd72272d7821adaf274ad91c288277713d9c18820d8499a7ff69e9a" +dependencies = [ + "typenum", + "version_check", +] + [[package]] name = "getrandom" version = "0.2.17" @@ -526,6 +590,15 @@ version = "0.5.0" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "2304e00983f87ffb38b55b444b5e3b60a884b5d30c0fca7d82fe33449bbe55ea" +[[package]] +name = "hmac" +version = "0.12.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "6c49c37c09c17a53d937dfbb742eb3a961d65a994e6bcdcf37e7399d0cc8ab5e" +dependencies = [ + "digest", +] + [[package]] name = "http" version = "1.4.2" @@ -602,11 +675,11 @@ dependencies = [ "http", "hyper", "hyper-util", - "rustls", + "rustls 0.23.41", "tokio", - "tokio-rustls", + "tokio-rustls 0.26.4", "tower-service", - "webpki-roots", + "webpki-roots 1.0.8", ] [[package]] @@ -1075,7 +1148,7 @@ dependencies = [ "quinn-proto", "quinn-udp", "rustc-hash", - "rustls", + "rustls 0.23.41", "socket2", "thiserror 2.0.18", "tokio", @@ -1096,7 +1169,7 @@ dependencies = [ "rand_pcg", "ring", "rustc-hash", - "rustls", + "rustls 0.23.41", "rustls-pki-types", "slab", "thiserror 2.0.18", @@ -1285,14 +1358,14 @@ dependencies = [ "percent-encoding", "pin-project-lite", "quinn", - "rustls", + "rustls 0.23.41", "rustls-pki-types", "serde", "serde_json", "serde_urlencoded", "sync_wrapper", "tokio", - "tokio-rustls", + "tokio-rustls 0.26.4", "tokio-util", "tower", "tower-http", @@ -1302,7 +1375,7 @@ dependencies = [ "wasm-bindgen-futures", "wasm-streams", "web-sys", - "webpki-roots", + "webpki-roots 1.0.8", ] [[package]] @@ -1338,6 +1411,20 @@ dependencies = [ "windows-sys 0.61.2", ] +[[package]] +name = "rustls" +version = "0.22.4" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "bf4ef73721ac7bcd79b2b315da7779d8fc09718c6b3d2d1b2d94850eb8c18432" +dependencies = [ + "log", + "ring", + "rustls-pki-types", + "rustls-webpki 0.102.8", + "subtle", + "zeroize", +] + [[package]] name = "rustls" version = "0.23.41" @@ -1347,7 +1434,7 @@ dependencies = [ "once_cell", "ring", "rustls-pki-types", - "rustls-webpki", + "rustls-webpki 0.103.13", "subtle", "zeroize", ] @@ -1362,6 +1449,17 @@ dependencies = [ "zeroize", ] +[[package]] +name = "rustls-webpki" +version = "0.102.8" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "64ca1bc8749bd4cf37b5ce386cc146580777b4e8572c7b97baf22c83f444bee9" +dependencies = [ + "ring", + "rustls-pki-types", + "untrusted", +] + [[package]] name = "rustls-webpki" version = "0.103.13" @@ -1461,6 +1559,28 @@ dependencies = [ "serde", ] +[[package]] +name = "sha1" +version = "0.10.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "e3bf829a2d51ab4a5ddf1352d8470c140cadc8301b2ae1789db023f01cedd6ba" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + +[[package]] +name = "sha2" +version = "0.10.9" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "a7507d819769d01a365ab707794a4084392c824f54a7a6a7862f8c3d0892b283" +dependencies = [ + "cfg-if", + "cpufeatures 0.2.17", + "digest", +] + [[package]] name = "shlex" version = "2.0.1" @@ -1666,14 +1786,41 @@ dependencies = [ "syn", ] +[[package]] +name = "tokio-rustls" +version = "0.25.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "775e0c0f0adb3a2f22a00c4745d728b479985fc15ee7ca6a2608388c5569860f" +dependencies = [ + "rustls 0.22.4", + "rustls-pki-types", + "tokio", +] + [[package]] name = "tokio-rustls" version = "0.26.4" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "1729aa945f29d91ba541258c8df89027d5792d85a8841fb65e8bf0f4ede4ef61" dependencies = [ - "rustls", + "rustls 0.23.41", + "tokio", +] + +[[package]] +name = "tokio-tungstenite" +version = "0.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "c83b561d025642014097b66e6c1bb422783339e0909e4429cde4749d1990bc38" +dependencies = [ + "futures-util", + "log", + "rustls 0.22.4", + "rustls-pki-types", "tokio", + "tokio-rustls 0.25.0", + "tungstenite", + "webpki-roots 0.26.11", ] [[package]] @@ -1810,6 +1957,33 @@ version = "0.2.5" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "e421abadd41a4225275504ea4d6566923418b7f05506fbc9c0fe86ba7396114b" +[[package]] +name = "tungstenite" +version = "0.21.0" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "9ef1a641ea34f399a848dea702823bbecfb4c486f911735368f1f137cb8257e1" +dependencies = [ + "byteorder", + "bytes", + "data-encoding", + "http", + "httparse", + "log", + "rand 0.8.6", + "rustls 0.22.4", + "rustls-pki-types", + "sha1", + "thiserror 1.0.69", + "url", + "utf-8", +] + +[[package]] +name = "typenum" +version = "1.20.1" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "b6f5e870be6c3b371b77fe0ee0bafb859fa4964b4404c27de1d380043c4dda20" + [[package]] name = "unicode-ident" version = "1.0.24" @@ -1840,6 +2014,12 @@ dependencies = [ "serde", ] +[[package]] +name = "utf-8" +version = "0.7.6" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "09cc8ee72d2a9becf2f2febe0205bbed8fc6615b7cb429ad062dc7b7ddd036a9" + [[package]] name = "utf8_iter" version = "1.0.4" @@ -1852,6 +2032,12 @@ version = "0.2.2" source = "registry+https://github.com/rust-lang/crates.io-index" checksum = "06abde3611657adf66d383f00b093d7faecc7fa57071cce2578660c9f1010821" +[[package]] +name = "version_check" +version = "0.9.5" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "0b928f33d975fc6ad9f86c8f283853ad26bdd5b10b7f1542aa2fa15e2289105a" + [[package]] name = "want" version = "0.3.1" @@ -1964,6 +2150,15 @@ dependencies = [ "wasm-bindgen", ] +[[package]] +name = "webpki-roots" +version = "0.26.11" +source = "registry+https://github.com/rust-lang/crates.io-index" +checksum = "521bc38abb08001b01866da9f51eb7c5d647a19260e00054a8c7fd5f9e57f7a9" +dependencies = [ + "webpki-roots 1.0.8", +] + [[package]] name = "webpki-roots" version = "1.0.8" diff --git a/crates/aibridge-core/Cargo.toml b/crates/aibridge-core/Cargo.toml index f8cb667..88632c0 100644 --- a/crates/aibridge-core/Cargo.toml +++ b/crates/aibridge-core/Cargo.toml @@ -24,6 +24,11 @@ bytes.workspace = true base64 = "0.22" # 随机数(router.rs 的 round_robin/random/weighted 策略用) rand = "0.8" +# Edge TTS WebSocket 协议(阶段2c edge-tts 适配器用) +tokio-tungstenite = { version = "0.21", features = ["connect", "rustls-tls-webpki-roots"] } +# HMAC-SHA256(Edge TTS 的 Sec-MS-GEC token 计算) +hmac = "0.12" +sha2 = "0.10" [dev-dependencies] # HTTP mock,用于适配器单测(openai_compat 等) diff --git a/crates/aibridge-core/src/adapters/edge_tts.rs b/crates/aibridge-core/src/adapters/edge_tts.rs new file mode 100644 index 0000000..44c7b3a --- /dev/null +++ b/crates/aibridge-core/src/adapters/edge_tts.rs @@ -0,0 +1,1669 @@ +//! Edge TTS 适配器(免费神经语音合成,免认证) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/audio_adapters.py` 的 `EdgeTTSAdapter`。 +//! +//! Edge TTS 基于微软 Edge 浏览器的免费神经语音合成服务,底层是 Azure 神经语音引擎, +//! 但**不需要 API Key**,完全免费。支持 100+ 种语音,覆盖 50+ 语言,中文支持优秀。 +//! +//! ## 协议 +//! +//! Edge TTS 的 `list_voices` 是普通 HTTP GET(返回 JSON 音色列表), +//! 而 `speech` 合成走 WebSocket(SSML 输入,二进制音频输出)。 +//! 两者都需要 `Sec-MS-GEC` token(基于固定 `TrustedClientToken` + 当前时间戳的 HMAC-SHA256)。 +//! +//! - list_voices: `GET https://speech.platform.bing.com/.../voices/list?TrustedClientToken=...&Sec-MS-GEC=...` +//! - speech: `wss://speech.platform.bing.com/.../edge/v1?TrustedClientToken=...&Sec-MS-GEC=...&ConnectionId=...` +//! +//! ## 特性(v1.3.3 保留) +//! +//! - `requires_api_key = false`(免认证) +//! - 音色自动降级:`voice` 传候选列表时,某音色失败(VoiceNotAvailable / ServiceUnavailable) +//! 自动切换下一个,直到成功或全部失败 +//! - 空音频检测:合成返回空音频时查 `list_voices` 区分语义——voice 在线则判服务端临时不可用 +//! (可重试),voice 不在线则判已下线(应换音色) +//! - 音色列表缓存(避免每次空音频都网络查询) + +use std::collections::HashMap; +use std::future::Future; +use std::time::{SystemTime, UNIX_EPOCH}; + +use async_trait::async_trait; +use futures::{SinkExt, StreamExt}; +use hmac::{Hmac, Mac}; +use rand::Rng; +use serde::Deserialize; +use serde_json::Value; +use sha2::Sha256; +use tokio::sync::Mutex; +use tokio_tungstenite::tungstenite::client::IntoClientRequest; +use tokio_tungstenite::tungstenite::http::{HeaderName, HeaderValue}; +use tokio_tungstenite::tungstenite::Message; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::audio::{SpeechRequest, SpeechResult}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; + +// ==================== 常量 ==================== + +/// Provider 类型标识 +const PROVIDER_TYPE: &str = "edge-tts"; +/// Provider 显示名称 +const PROVIDER_NAME: &str = "Edge TTS"; +/// Edge TTS 固定的可信客户端 token(社区逆向所得,所有客户端共用) +const TRUSTED_CLIENT_TOKEN: &str = "6A5AA1D4EAFF4E9FB37E23D68491D6F4"; +/// Sec-MS-GEC 版本号(对应 Edge 浏览器版本) +const SEC_MS_GEC_VERSION: &str = "1-130.0.2849.68"; +/// 默认音色(中文女声晓晓) +const DEFAULT_VOICE: &str = "zh-CN-XiaoxiaoNeural"; +/// 默认 API 基地址 +const DEFAULT_API_BASE: &str = "https://speech.platform.bing.com"; +/// 音色列表端点路径 +const VOICES_PATH: &str = "/consumer/speech/synthesize/readaloud/voices/list"; +/// WebSocket 合成端点路径 +const WS_PATH: &str = "/consumer/speech/synthesize/readaloud/edge/v1"; +/// WS 握手所需 Origin header +const WS_ORIGIN: &str = "https://speech.platform.bing.com"; +/// WS 握手 User-Agent +const WS_USER_AGENT: &str = "Mozilla/5.0 (Windows NT 10.0; Win64; x64)"; +/// Windows epoch 偏移(1601-01-01 到 1970-01-01 的 100ns 单位) +const WIN_EPOCH: i64 = 116_444_736_000_000_000; +/// 5 分钟对应的 100ns 单位(GEC token 向下取整到 5 分钟边界) +const GEC_WINDOW_TICKS: i64 = 3_000_000_000; + +/// HMAC-SHA256 类型别名(计算 Sec-MS-GEC token) +type HmacSha256 = Hmac; + +// ==================== 音色别名查询 ==================== + +/// 查询常见音色别名(简称/中文名)→ 完整 voice ID +/// +/// 对应 Python v1 `EdgeTTSAdapter.COMMON_VOICES`。返回 None 表示非已知别名。 +fn lookup_common_voice(voice: &str) -> Option<&'static str> { + let v = match voice { + // 中文女声 + "xiaoxiao" | "晓晓" => "zh-CN-XiaoxiaoNeural", + "xiaoyi" | "晓伊" => "zh-CN-XiaoyiNeural", + "xiaochen" | "晓辰" => "zh-CN-XiaochenNeural", + "xiaohan" | "晓涵" => "zh-CN-XiaohanNeural", + "xiaomeng" | "晓梦" => "zh-CN-XiaomengNeural", + "xiaomo" | "晓墨" => "zh-CN-XiaomoNeural", + "xiaoqiu" | "晓秋" => "zh-CN-XiaoqiuNeural", + "xiaorui" | "晓睿" => "zh-CN-XiaoruiNeural", + "xiaoshuang" | "晓双" => "zh-CN-XiaoshuangNeural", + "xiaoxuan" | "晓萱" => "zh-CN-XiaoxuanNeural", + "xiaoyan" | "晓颜" => "zh-CN-XiaoyanNeural", + "xiaoyou" | "晓悠" => "zh-CN-XiaoyouNeural", + // 中文男声 + "yunjian" | "云健" => "zh-CN-YunjianNeural", + "yunxi" | "云希" => "zh-CN-YunxiNeural", + "yunxia" | "云夏" => "zh-CN-YunxiaNeural", + "yunyang" | "云扬" => "zh-CN-YunyangNeural", + "yunque" | "云泽" => "zh-CN-YunzeNeural", + // 英文女声 + "jenny" => "en-US-JennyNeural", + "jenny-multilingual" => "en-US-JennyMultilingualNeural", + "aria" => "en-US-AriaNeural", + // 英文男声 + "guy" => "en-US-GuyNeural", + "roger" => "en-US-RogerNeural", + "davis" => "en-US-DavisNeural", + "tony" => "en-US-TonyNeural", + "jason" => "en-US-JasonNeural", + // 日文 + "nanami" => "ja-JP-NanamiNeural", + "keita" => "ja-JP-KeitaNeural", + // 韩文 + "sun-hi" => "ko-KR-SunHiNeural", + "in-jun" => "ko-KR-InJoonNeural", + // 法文 + "denise" => "fr-FR-DeniseNeural", + "henri" => "fr-FR-HenriNeural", + // 德文 + "katja" => "de-DE-KatjaNeural", + "conrad" => "de-DE-ConradNeural", + // 西班牙文 + "elvira" => "es-ES-ElviraNeural", + "alvaro" => "es-ES-AlvaroNeural", + _ => return None, + }; + Some(v) +} + +// ==================== Edge 原始音色反序列化结构 ==================== + +/// Edge TTS list_voices 返回的单个音色项(原始 JSON 结构) +/// +/// 字段名与 Edge 服务端返回的 JSON 一致(PascalCase)。 +#[derive(Debug, Deserialize)] +struct EdgeVoiceRaw { + #[serde(rename = "ShortName")] + short_name: String, + #[serde(rename = "Name")] + name: String, + #[serde(rename = "Locale")] + locale: String, + #[serde(rename = "Gender")] + gender: String, + #[serde(rename = "FriendlyName", default)] + friendly_name: Option, +} + +// ==================== EdgeTtsAdapter ==================== + +/// Edge TTS 适配器 +/// +/// 持有 HTTP 客户端(list_voices 用)与音色列表缓存。 +/// speech 合成按需建立 WebSocket 连接(per-request,不持有长连接)。 +pub struct EdgeTtsAdapter { + /// Provider 配置 + #[allow(dead_code)] + config: ProviderConfig, + /// HTTP 客户端(list_voices 的 HTTP GET) + http: HttpClient, + /// 音色列表缓存(避免重复网络拉取) + voices_cache: Mutex>>, + /// API 基地址(默认 `https://speech.platform.bing.com`,可由 config.base_url 覆盖) + api_base: String, +} + +impl EdgeTtsAdapter { + /// 创建 Edge TTS 适配器 + /// + /// `config.base_url` 可覆盖 API 基地址(主要用于测试指向 mock server), + /// 为空时用 `DEFAULT_API_BASE`。Edge TTS 免认证,不强制 api_key。 + pub fn new(config: ProviderConfig) -> Result { + let api_base = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_BASE.to_string()); + let opts = ClientOptions::builder().timeout(config.timeout).build(); + let http = HttpClient::new(&opts)?; + Ok(Self { + config, + http, + voices_cache: Mutex::new(None), + api_base, + }) + } + + // ==================== 纯函数(协议构造/解析,可单测) ==================== + + /// 解析音色名:别名/中文名/完整 ID → 完整 voice ID + /// + /// 对应 Python v1 `EdgeTTSAdapter._resolve_voice`。 + /// - 空值 → 默认音色 + /// - 已知别名(精确或大小写不敏感)→ 对应完整 ID + /// - 含 ≥2 个 `-`(如 `zh-CN-XiaoxiaoNeural`)→ 视为完整 ID 原样返回 + /// - 其余 → 默认音色 + fn resolve_voice(voice: &str) -> String { + if voice.is_empty() { + return DEFAULT_VOICE.to_string(); + } + if let Some(v) = lookup_common_voice(voice) { + return v.to_string(); + } + // 大小写不敏感匹配(中文别名 lowercase 无影响,英文如 Xiaoxiao → xiaoxiao) + if let Some(v) = lookup_common_voice(&voice.to_lowercase()) { + return v.to_string(); + } + // 完整 voice ID:含至少 2 个 `-`(语言-区域-名称) + if voice.contains('-') && voice.matches('-').count() >= 2 { + return voice.to_string(); + } + DEFAULT_VOICE.to_string() + } + + /// 获取输出格式:返回 (edge-tts format string, content_type, 扩展名) + /// + /// 对应 Python v1 `EdgeTTSAdapter._get_output_format`。 + fn get_output_format(fmt: Option<&str>) -> (&'static str, &'static str, &'static str) { + let key = fmt.unwrap_or("mp3").to_lowercase(); + let edge_fmt: &str = match key.as_str() { + "mp3" => "audio-24khz-48kbit-mp3-mono", + "mp3-96k" => "audio-24khz-96kbit-mp3-mono", + "mp3-128k" => "audio-24khz-128kbit-mp3-mono", + "mp3-160k" => "audio-24khz-160kbit-mp3-mono", + "webm" => "audio-24khz-48kbit-opus-mono", + "webm-24khz-16bit-mono-opus" => "audio-24khz-16bit-mono-opus", + "ogg" => "audio-24khz-48kbit-opus-mono", + "wav" | "pcm" => "audio-24khz-16bit-mono-pcm", + _ => "audio-24khz-48kbit-mp3-mono", + }; + let (content_type, ext) = if edge_fmt.contains("mp3") { + ("audio/mpeg", "mp3") + } else if edge_fmt.contains("opus") || edge_fmt.contains("webm") || edge_fmt.contains("ogg") + { + ("audio/ogg", "ogg") + } else if edge_fmt.contains("pcm") || edge_fmt.contains("wav") { + ("audio/wav", "wav") + } else { + ("audio/mpeg", "mp3") + }; + (edge_fmt, content_type, ext) + } + + /// 计算 Sec-MS-GEC token(HMAC-SHA256) + /// + /// 算法(社区逆向):把当前 Unix 秒转为 Windows ticks(100ns 单位), + /// 向下取整到 5 分钟边界,用 `TRUSTED_CLIENT_TOKEN` 作 key 做 HMAC-SHA256, + /// 输出大写十六进制。服务端用此 token 鉴权(无需 API Key)。 + fn compute_sec_ms_gec(unix_secs: i64) -> String { + let mut ticks = unix_secs * 10_000_000 + WIN_EPOCH; + ticks -= ticks.rem_euclid(GEC_WINDOW_TICKS); + let mut mac = HmacSha256::new_from_slice(TRUSTED_CLIENT_TOKEN.as_bytes()) + .expect("HMAC 接受任意长度 key"); + mac.update(&ticks.to_be_bytes()); + let result = mac.finalize().into_bytes(); + let mut hex = String::with_capacity(64); + for b in result.iter() { + hex.push_str(&format!("{:02X}", b)); + } + hex + } + + /// 构造 WebSocket 合成端点 URL + /// + /// `api_base` 的 `https://` 替换为 `wss://`(`http://` → `ws://`)。 + fn build_ws_url(gec: &str, connection_id: &str, api_base: &str) -> String { + let ws_base = api_base + .replacen("https://", "wss://", 1) + .replacen("http://", "ws://", 1); + format!( + "{base}{path}?TrustedClientToken={token}&Sec-MS-GEC={gec}&Sec-MS-GEC-Version={ver}&ConnectionId={cid}", + base = ws_base.trim_end_matches('/'), + path = WS_PATH, + token = TRUSTED_CLIENT_TOKEN, + gec = gec, + ver = SEC_MS_GEC_VERSION, + cid = connection_id, + ) + } + + /// 构造 list_voices HTTP 端点 URL + fn build_voices_url(gec: &str, api_base: &str) -> String { + format!( + "{base}{path}?TrustedClientToken={token}&Sec-MS-GEC={gec}&Sec-MS-GEC-Version={ver}", + base = api_base.trim_end_matches('/'), + path = VOICES_PATH, + token = TRUSTED_CLIENT_TOKEN, + gec = gec, + ver = SEC_MS_GEC_VERSION, + ) + } + + /// 构造 WebSocket 握手请求(设置 Origin / User-Agent header) + /// + /// Edge TTS 服务端要求 Origin header,否则拒绝握手。 + fn build_ws_request(url: &str) -> Result> { + let mut request = url + .into_client_request() + .map_err(|e| AibridgeError::validation(format!("Edge TTS WebSocket URL 无效: {e}")))?; + let headers = request.headers_mut(); + headers.insert( + HeaderName::from_static("origin"), + HeaderValue::from_static(WS_ORIGIN), + ); + headers.insert( + HeaderName::from_static("user-agent"), + HeaderValue::from_static(WS_USER_AGENT), + ); + Ok(request) + } + + /// 从 voice ID 提取 locale(如 `zh-CN-XiaoxiaoNeural` → `zh-CN`) + fn locale_from_voice(voice: &str) -> String { + let parts: Vec<&str> = voice.splitn(3, '-').collect(); + if parts.len() >= 2 { + format!("{}-{}", parts[0], parts[1]) + } else { + "en-US".to_string() + } + } + + /// 构造 SSML(语音合成标记语言) + /// + /// 对应 Python v1 构造的 `...` 结构。 + fn build_ssml(text: &str, voice: &str, rate: &str, pitch: &str, volume: &str) -> String { + let locale = Self::locale_from_voice(voice); + let escaped = escape_xml(text); + format!( + "\ + \ + \ + {escaped}\ + " + ) + } + + /// 构造 WS 配置消息(speech.config,声明输出格式) + fn build_config_message(output_format: &str, timestamp: &str) -> String { + let body = serde_json::json!({ + "context": { + "synthesis": { + "audio": { + "metadataoptions": { + "sentenceBoundaryEnabled": "false", + "wordBoundaryEnabled": "false" + }, + "outputFormat": output_format + } + } + } + }); + format!( + "X-Timestamp:{ts}\r\n\ + Content-Type:application/json; charset=utf-8\r\n\ + Path:speech.config\r\n\r\n\ + {body}", + ts = timestamp, + body = body, + ) + } + + /// 构造 WS 合成消息(ssml,携带待合成文本) + fn build_synth_message(ssml: &str, request_id: &str, timestamp: &str) -> String { + format!( + "X-RequestId:{rid}\r\n\ + Content-Type:application/ssml+xml\r\n\ + X-Timestamp:{ts}\r\n\ + Path:ssml\r\n\r\n\ + {ssml}", + rid = request_id, + ts = timestamp, + ssml = ssml, + ) + } + + /// 解析 WS 二进制消息,提取音频数据 + /// + /// Edge TTS binary 消息格式:2 字节大端 header 长度 + header 文本 + audio 数据。 + /// 仅当 header 含 `Path:audio` 时返回音频切片,其余返回 None。 + fn parse_audio_binary(msg: &[u8]) -> Option<&[u8]> { + if msg.len() < 2 { + return None; + } + let header_len = u16::from_be_bytes([msg[0], msg[1]]) as usize; + if header_len == 0 || msg.len() < 2 + header_len { + return None; + } + let header = &msg[2..2 + header_len]; + let header_str = std::str::from_utf8(header).ok()?; + if header_str.contains("Path:audio") { + Some(&msg[2 + header_len..]) + } else { + None + } + } + + /// 解析 WS 文本消息,提取 `Path:` 字段值(如 `turn.end` / `response`) + fn parse_ws_text_path(msg: &str) -> Option<&str> { + for line in msg.lines() { + if let Some(rest) = line.strip_prefix("Path:") { + return Some(rest.trim()); + } + } + None + } + + /// 从 WS 文本消息中提取错误(`Path:response` 且 body `"Type":"error"`) + /// + /// 音色相关错误(Code/Message 含 "voice")→ VoiceNotAvailable;其余 → Api。 + fn extract_ws_error(msg: &str) -> Option { + let body = msg + .split("\r\n\r\n") + .nth(1) + .or_else(|| msg.split("\n\n").nth(1))?; + let v: Value = serde_json::from_str(body.trim()).ok()?; + if v.get("Type").and_then(|t| t.as_str()) != Some("error") { + return None; + } + let code = v.get("Code").and_then(|c| c.as_str()).unwrap_or("UNKNOWN"); + let message = v + .get("Message") + .and_then(|m| m.as_str()) + .unwrap_or("Edge TTS 合成错误"); + let lower = format!("{code} {message}").to_lowercase(); + if lower.contains("voice") { + Some(AibridgeError::voice_not_available(format!( + "{code}: {message}" + ))) + } else { + Some(AibridgeError::api(0, format!("{code}: {message}"))) + } + } + + /// 空音频错误分类(纯函数,便于单测) + /// + /// voice 在线但返空 → ServiceUnavailable(可重试); + /// voice 不在线 → VoiceNotAvailable(应换音色)。 + fn empty_audio_error(voice_id: &str, voice_available: bool) -> AibridgeError { + if voice_available { + AibridgeError::service_unavailable(format!( + "Edge TTS 服务端返回空音频(voice={} 仍在线),可能是限流或网络抖动,可重试", + voice_id + )) + } else { + AibridgeError::voice_not_available(format!( + "Edge TTS 音色 {} 已下线或不存在,请更换音色", + voice_id + )) + } + } + + /// 构造语速字符串(edge-tts 格式 `+0%` / `-50%` / `+100%`) + /// + /// 优先用 `extra.rate` 字符串(与 Python 兼容),否则从 `speed` 数值换算 + /// (speed=1.0 → `+0%`,speed=0.5 → `-50%`,speed=2.0 → `+100%`)。 + fn build_rate(speed: Option, extra: &HashMap) -> String { + if let Some(r) = extra.get("rate").and_then(|v| v.as_str()) { + return r.to_string(); + } + let speed = speed.unwrap_or(1.0); + let pct = ((speed - 1.0) * 100.0).round() as i32; + format!("{pct:+}%") + } + + /// 构造音调字符串(edge-tts 格式 `+0Hz` / `+100Hz` / `-50Hz`) + /// + /// 优先用 `extra.pitch` 字符串,否则从 `pitch` 数值换算(pitch∈[-1,1] → ±100Hz)。 + fn build_pitch(pitch: Option, extra: &HashMap) -> String { + if let Some(p) = extra.get("pitch").and_then(|v| v.as_str()) { + return p.to_string(); + } + let pitch = pitch.unwrap_or(0.0); + let hz = (pitch * 100.0).round() as i32; + format!("{hz:+}Hz") + } + + /// 构造音量字符串(edge-tts 格式 `+0%` / `-50%` / `+100%`) + /// + /// 优先用 `extra.volume` 字符串,否则从 `volume` 数值换算 + /// (volume=1.0 → `+0%`,volume=0.5 → `-50%`,volume=2.0 → `+100%`)。 + fn build_volume(volume: Option, extra: &HashMap) -> String { + if let Some(v) = extra.get("volume").and_then(|v| v.as_str()) { + return v.to_string(); + } + let volume = volume.unwrap_or(1.0); + let pct = ((volume - 1.0) * 100.0).round() as i32; + format!("{pct:+}%") + } + + /// 把 Unix 秒格式化为 ISO 8601 UTC 时间戳(`2023-11-14T22:13:20.000Z`) + /// + /// 用 Howard Hinnant 的 civil_from_days 算法,避免引入 chrono 依赖。 + fn utc_timestamp_iso(unix_secs: i64) -> String { + let days = unix_secs.div_euclid(86400); + let secs_of_day = unix_secs.rem_euclid(86400); + let hour = secs_of_day / 3600; + let min = (secs_of_day % 3600) / 60; + let sec = secs_of_day % 60; + + // civil_from_days(Howard Hinnant) + let z = days + 719468; + let era = (if z >= 0 { z } else { z - 146096 }) / 146097; + let doe = z - era * 146097; + let yoe = (doe - doe / 1460 + doe / 36524 - doe / 146096) / 365; + let y = yoe + era * 400; + let doy = doe - (365 * yoe + yoe / 4 - yoe / 100); + let mp = (5 * doy + 2) / 153; + let d = doy - (153 * mp + 2) / 5 + 1; + let m = if mp < 10 { mp + 3 } else { mp - 9 }; + let year = if m <= 2 { y + 1 } else { y }; + + format!( + "{:04}-{:02}-{:02}T{:02}:{:02}:{:02}.000Z", + year, m, d, hour, min, sec + ) + } + + /// 判断错误是否可触发音色降级(VoiceNotAvailable / ServiceUnavailable) + fn is_fallback_error(e: &AibridgeError) -> bool { + matches!( + e, + AibridgeError::VoiceNotAvailable { .. } | AibridgeError::ServiceUnavailable { .. } + ) + } + + /// 把 Edge 原始音色转为统一 `VoiceInfo` + fn convert_voice(raw: EdgeVoiceRaw) -> VoiceInfo { + let display_name = raw.friendly_name.unwrap_or(raw.name); + VoiceInfo::builder() + .short_name(raw.short_name.clone()) + .name(display_name) + .locale(raw.locale) + .gender(raw.gender) + .voice_id(raw.short_name) + .build() + } + + // ==================== 音色降级(可单测) ==================== + + /// 音色降级合成:按候选列表逐个尝试,失败(VoiceNotAvailable / ServiceUnavailable) + /// 切换下一个,直到成功或全部失败。其余错误不降级直接抛出。 + /// + /// 抽成独立函数便于单测降级逻辑(`single` 闭包模拟单音色合成结果)。 + async fn speech_with_fallback( + candidates: Vec, + mut single: F, + ) -> Result + where + F: FnMut(&str) -> Fut + Send, + Fut: Future> + Send, + { + let total = candidates.len(); + let mut last_error: Option = None; + for (idx, voice_raw) in candidates.iter().enumerate() { + match single(voice_raw).await { + Ok(result) => return Ok(result), + Err(e) if Self::is_fallback_error(&e) => { + // 单 voice 模式:直接抛当前异常,无需存 last_error + if total <= 1 { + return Err(e); + } + // 多 voice 列表模式:记录最后一个失败异常,尝试下一个 + last_error = Some(e); + tracing::warn!( + "Edge TTS voice {} 合成失败,尝试下一个候选 ({}/{})", + voice_raw, + idx + 1, + total + ); + continue; + } + Err(e) => return Err(e), + } + } + Err(last_error + .unwrap_or_else(|| AibridgeError::api(0, "Edge TTS 语音合成失败,无可用 voice"))) + } + + // ==================== 网络方法 ==================== + + /// 单音色合成(WebSocket 协议) + /// + /// 建连 → 发配置消息 → 发 SSML 合成消息 → 收集 audio binary → 检测空音频。 + /// 不含降级逻辑(由 `speech` 包装降级)。 + #[allow(clippy::too_many_arguments)] // 合成参数均为必需,无法再精简 + async fn speech_single( + &self, + model: &str, + input: &str, + voice_id: &str, + output_format: &str, + rate: &str, + pitch: &str, + volume: &str, + ) -> Result { + let (edge_fmt, content_type, ext) = Self::get_output_format(Some(output_format)); + + let unix_secs = current_unix_secs(); + let gec = Self::compute_sec_ms_gec(unix_secs); + let conn_id = generate_connection_id(); + let url = Self::build_ws_url(&gec, &conn_id, &self.api_base); + let request = Self::build_ws_request(&url)?; + + let (ws_stream, _) = tokio_tungstenite::connect_async(request) + .await + .map_err(map_ws_error)?; + let mut ws = ws_stream; + + let ts = Self::utc_timestamp_iso(unix_secs); + // 发送配置消息(声明输出格式) + let config_msg = Self::build_config_message(edge_fmt, &ts); + ws.send(Message::Text(config_msg)) + .await + .map_err(map_ws_error)?; + + // 发送合成消息(SSML) + let request_id = generate_connection_id(); + let ssml = Self::build_ssml(input, voice_id, rate, pitch, volume); + let synth_msg = Self::build_synth_message(&ssml, &request_id, &ts); + ws.send(Message::Text(synth_msg)) + .await + .map_err(map_ws_error)?; + + // 接收消息,收集音频二进制 + let mut audio_chunks: Vec> = Vec::new(); + while let Some(msg_result) = ws.next().await { + let msg = msg_result.map_err(map_ws_error)?; + match msg { + Message::Binary(data) => { + if let Some(audio) = Self::parse_audio_binary(&data) { + audio_chunks.push(audio.to_vec()); + } + } + Message::Text(text) => { + if let Some(path) = Self::parse_ws_text_path(&text) { + match path { + "turn.end" => break, + "response" => { + if let Some(err) = Self::extract_ws_error(&text) { + return Err(err); + } + } + _ => {} + } + } + } + Message::Close(_) => break, + _ => {} + } + } + + let audio_data: Vec = audio_chunks.into_iter().flatten().collect(); + + // 空音频检测:区分 voice 在线/下线 + if audio_data.is_empty() { + let available = self.check_voice_available(voice_id).await; + return Err(Self::empty_audio_error(voice_id, available)); + } + + Ok(SpeechResult { + audio_data: Some(audio_data), + audio_url: None, + audio_base64: None, + content_type: content_type.to_string(), + format: ext.to_string(), + duration: None, + model: Some(if model.is_empty() { + PROVIDER_TYPE.to_string() + } else { + model.to_string() + }), + }) + } + + /// 拉取全量音色列表(带缓存) + /// + /// 缓存未命中时 HTTP GET 服务端,缓存后直接返回(language 过滤在 `list_voices` 做)。 + async fn list_voices_raw(&self) -> Result> { + // 先检查缓存 + { + let cache = self.voices_cache.lock().await; + if let Some(ref voices) = *cache { + return Ok(voices.clone()); + } + } + let unix_secs = current_unix_secs(); + let gec = Self::compute_sec_ms_gec(unix_secs); + let url = Self::build_voices_url(&gec, &self.api_base); + + let resp = self.http.inner().get(&url).send().await.map_err(|e| { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } + })?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_edge_http_error(status_code, &body)); + } + let raw_voices: Vec = resp.json().await.map_err(AibridgeError::from)?; + let voices: Vec = raw_voices.into_iter().map(Self::convert_voice).collect(); + + // 写缓存 + let mut cache = self.voices_cache.lock().await; + *cache = Some(voices.clone()); + Ok(voices) + } + + /// 检查指定 voice 是否在可用音色列表中(空音频时区分异常语义) + /// + /// list_voices 查询本身失败时保守返回 true(按服务端临时问题处理)。 + async fn check_voice_available(&self, voice_id: &str) -> bool { + match self.list_voices(None).await { + Ok(voices) => voices + .iter() + .any(|v| v.short_name.as_deref() == Some(voice_id)), + Err(_) => true, + } + } +} + +#[async_trait] +impl Adapter for EdgeTtsAdapter { + fn provider_type(&self) -> &str { + PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::AudioSpeech); + caps.insert(Capabilities::ListVoices); + caps + } + + /// 免认证:Edge TTS 无需 API Key + fn requires_api_key(&self) -> bool { + false + } + + async fn start(&mut self) -> Result<()> { + // 无需惰性初始化(依赖编译期链接),保持 no-op + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // WS 连接是 per-request 的,无需释放长连接资源 + Ok(()) + } + + /// 文字转语音(含音色自动降级) + async fn speech(&self, req: SpeechRequest) -> Result { + // 规范化为候选列表:空列表兜底为 [""](走默认音色) + let candidates: Vec = if req.voice.voices.is_empty() { + vec![String::new()] + } else { + req.voice.voices.clone() + }; + let model = req.model.clone(); + let input = req.input.clone(); + let fmt = req.response_format.clone(); + let rate = Self::build_rate(req.speed, &req.extra); + let pitch = Self::build_pitch(req.pitch, &req.extra); + let volume = Self::build_volume(req.volume, &req.extra); + + let single = move |voice_raw: &str| { + let voice_id = Self::resolve_voice(voice_raw); + let model = model.clone(); + let input = input.clone(); + let fmt = fmt.clone(); + let rate = rate.clone(); + let pitch = pitch.clone(); + let volume = volume.clone(); + async move { + self.speech_single(&model, &input, &voice_id, &fmt, &rate, &pitch, &volume) + .await + } + }; + Self::speech_with_fallback(candidates, single).await + } + + /// 列出可用音色(带缓存,按 language 前缀过滤) + async fn list_voices(&self, language: Option<&str>) -> Result> { + let voices = self.list_voices_raw().await?; + match language { + Some(lang) if !lang.is_empty() => Ok(voices + .into_iter() + .filter(|v| { + v.locale + .as_deref() + .map(|l| l.starts_with(lang)) + .unwrap_or(false) + }) + .collect()), + _ => Ok(voices), + } + } + + /// 列出 Edge TTS 模型(无标准 /models 端点,保留硬编码列表) + async fn list_models(&self, filter: Option) -> Result> { + let models = vec![ModelInfo { + id: PROVIDER_TYPE.into(), + name: "Edge TTS Neural".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: false, + description: Some( + "微软 Edge 浏览器免费神经语音合成,支持 100+ 种语音,50+ 语言,中文支持优秀".into(), + ), + created: None, + }]; + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 错误映射 ==================== + +/// 将 WebSocket 错误映射为 AibridgeError +/// +/// 注意:`AibridgeError::Network` 仅接受 `reqwest::Error`,WS 错误无法转换, +/// 故 WS 网络错误映射到 `ServiceUnavailable`(语义:服务暂时不可达,可重试, +/// 与 Network 的 is_retryable 一致);WS 握手 HTTP 错误映射到 `Api`。 +fn map_ws_error(e: tokio_tungstenite::tungstenite::Error) -> AibridgeError { + use tokio_tungstenite::tungstenite::Error as WsErr; + match e { + WsErr::Http(resp) => { + let status = resp.status().as_u16(); + AibridgeError::api( + status, + format!("Edge TTS WebSocket 握手失败 (HTTP {status})"), + ) + } + _ => AibridgeError::service_unavailable(format!("Edge TTS WebSocket 网络错误: {e}")), + } +} + +/// 将 Edge TTS list_voices HTTP 错误映射为 AibridgeError +fn map_edge_http_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 | 403 => AibridgeError::authentication(format!("Edge TTS 认证失败: {body}")), + 429 => AibridgeError::rate_limit(format!("Edge TTS 限流: {body}")), + 404 => AibridgeError::voice_not_available(format!("Edge TTS 音色端点不存在: {body}")), + s if s >= 500 => { + AibridgeError::service_unavailable(format!("Edge TTS 服务不可用 ({s}): {body}")) + } + s => AibridgeError::api(s, format!("Edge TTS HTTP {s}: {body}")), + } +} + +// ==================== 辅助函数 ==================== + +/// XML 特殊字符转义(SSML 文本内容用) +fn escape_xml(s: &str) -> String { + let mut out = String::with_capacity(s.len()); + for c in s.chars() { + match c { + '<' => out.push_str("<"), + '>' => out.push_str(">"), + '&' => out.push_str("&"), + _ => out.push(c), + } + } + out +} + +/// 当前 Unix 时间戳(秒) +fn current_unix_secs() -> i64 { + SystemTime::now() + .duration_since(UNIX_EPOCH) + .map(|d| d.as_secs() as i64) + .unwrap_or(0) +} + +/// 生成 32 字符十六进制连接 ID(模拟无连字符 UUID,Edge TTS 要求 hex 格式) +fn generate_connection_id() -> String { + let id: u128 = rand::thread_rng().gen(); + format!("{:032x}", id) +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::model::audio::TranscribeRequest; + use crate::model::chat::ChatRequest; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + use std::sync::atomic::{AtomicU32, Ordering}; + use std::sync::Arc; + + /// 构造测试用适配器(base_url 指向 mockito server,None 用默认基地址) + fn make_adapter(base_url: Option) -> EdgeTtsAdapter { + let mut opts = ClientOptions::builder(); + if let Some(u) = base_url { + opts = opts.base_url(u); + } + let config = ProviderConfig::from_options(PROVIDER_TYPE, opts.build()); + EdgeTtsAdapter::new(config).expect("构造 EdgeTtsAdapter 失败") + } + + // ============ 基本属性 ============ + + #[test] + fn requires_api_key_is_false() { + let adapter = make_adapter(None); + assert!(!adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_speech_and_list_voices() { + let adapter = make_adapter(None); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::AudioSpeech)); + assert!(caps.contains(&Capabilities::ListVoices)); + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn provider_type_and_name() { + let adapter = make_adapter(None); + assert_eq!(adapter.provider_type(), "edge-tts"); + assert_eq!(adapter.provider_name(), "Edge TTS"); + } + + #[tokio::test] + async fn list_models_returns_single_edge_tts() { + let adapter = make_adapter(None); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 1); + assert_eq!(models[0].id, "edge-tts"); + assert_eq!(models[0].model_type, ModelType::Audio); + assert_eq!(models[0].provider, "edge-tts"); + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let adapter = make_adapter(None); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 1); + let chat = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat.is_empty()); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter(None); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 不支持能力(默认实现) ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter(None); + let req = ChatRequest::builder("m", vec![]).build(); + assert!(matches!( + adapter.chat(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn image_generate_returns_unsupported() { + let adapter = make_adapter(None); + let req = ImageRequest::builder("m", "p").build(); + assert!(matches!( + adapter.image_generate(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter(None); + let req = VideoRequest::builder("m", "p").build(); + assert!(matches!( + adapter.video_create(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn embed_returns_unsupported() { + let adapter = make_adapter(None); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + assert!(matches!( + adapter.embed(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn transcribe_returns_unsupported() { + let adapter = make_adapter(None); + let req = TranscribeRequest::builder("m", FileInput::path("/tmp/a.mp3")).build(); + assert!(matches!( + adapter.transcribe(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + // ============ resolve_voice ============ + + #[test] + fn resolve_voice_empty_returns_default() { + assert_eq!(EdgeTtsAdapter::resolve_voice(""), DEFAULT_VOICE); + } + + #[test] + fn resolve_voice_alias_lookup() { + assert_eq!( + EdgeTtsAdapter::resolve_voice("xiaoxiao"), + "zh-CN-XiaoxiaoNeural" + ); + assert_eq!( + EdgeTtsAdapter::resolve_voice("晓晓"), + "zh-CN-XiaoxiaoNeural" + ); + assert_eq!(EdgeTtsAdapter::resolve_voice("yunxi"), "zh-CN-YunxiNeural"); + assert_eq!(EdgeTtsAdapter::resolve_voice("jenny"), "en-US-JennyNeural"); + assert_eq!( + EdgeTtsAdapter::resolve_voice("nanami"), + "ja-JP-NanamiNeural" + ); + } + + #[test] + fn resolve_voice_case_insensitive() { + assert_eq!( + EdgeTtsAdapter::resolve_voice("Xiaoxiao"), + "zh-CN-XiaoxiaoNeural" + ); + assert_eq!(EdgeTtsAdapter::resolve_voice("JENNY"), "en-US-JennyNeural"); + } + + #[test] + fn resolve_voice_passes_through_full_id() { + assert_eq!( + EdgeTtsAdapter::resolve_voice("zh-CN-XiaoxiaoNeural"), + "zh-CN-XiaoxiaoNeural" + ); + assert_eq!( + EdgeTtsAdapter::resolve_voice("en-US-JennyMultilingualNeural"), + "en-US-JennyMultilingualNeural" + ); + } + + #[test] + fn resolve_voice_unknown_short_returns_default() { + assert_eq!(EdgeTtsAdapter::resolve_voice("unknown"), DEFAULT_VOICE); + } + + // ============ get_output_format ============ + + #[test] + fn get_output_format_default_mp3() { + let (fmt, ct, ext) = EdgeTtsAdapter::get_output_format(None); + assert_eq!(fmt, "audio-24khz-48kbit-mp3-mono"); + assert_eq!(ct, "audio/mpeg"); + assert_eq!(ext, "mp3"); + } + + #[test] + fn get_output_format_wav() { + let (fmt, ct, ext) = EdgeTtsAdapter::get_output_format(Some("wav")); + assert_eq!(fmt, "audio-24khz-16bit-mono-pcm"); + assert_eq!(ct, "audio/wav"); + assert_eq!(ext, "wav"); + } + + #[test] + fn get_output_format_webm_is_ogg() { + let (_, ct, ext) = EdgeTtsAdapter::get_output_format(Some("webm")); + assert_eq!(ct, "audio/ogg"); + assert_eq!(ext, "ogg"); + } + + #[test] + fn get_output_format_unknown_defaults_mp3() { + let (fmt, _, _) = EdgeTtsAdapter::get_output_format(Some("xyz")); + assert_eq!(fmt, "audio-24khz-48kbit-mp3-mono"); + } + + // ============ compute_sec_ms_gec ============ + + #[test] + fn compute_sec_ms_gec_is_uppercase_hex_64() { + let gec = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_000); + assert_eq!(gec.len(), 64); + assert!(gec + .chars() + .all(|c| c.is_ascii_uppercase() || c.is_ascii_digit())); + } + + #[test] + fn compute_sec_ms_gec_deterministic() { + let a = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_000); + let b = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_000); + assert_eq!(a, b); + } + + #[test] + fn compute_sec_ms_gec_same_within_5min_window() { + // 1700000000 在窗口 [1699999800, 1700000100) 内 + let a = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_000); + let b = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_099); + assert_eq!(a, b); + } + + #[test] + fn compute_sec_ms_gec_changes_across_boundary() { + let a = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_000); + let b = EdgeTtsAdapter::compute_sec_ms_gec(1_700_000_100); + assert_ne!(a, b); + } + + // ============ URL 构造 ============ + + #[test] + fn build_ws_url_contains_all_params() { + let url = + EdgeTtsAdapter::build_ws_url("GEC123", "conn456", "https://speech.platform.bing.com"); + assert!(url.starts_with("wss://speech.platform.bing.com")); + assert!(url.contains(WS_PATH)); + assert!(url.contains("TrustedClientToken=6A5AA1D4EAFF4E9FB37E23D68491D6F4")); + assert!(url.contains("Sec-MS-GEC=GEC123")); + assert!(url.contains("Sec-MS-GEC-Version=1-130.0.2849.68")); + assert!(url.contains("ConnectionId=conn456")); + } + + #[test] + fn build_ws_url_converts_http_to_ws() { + let url = EdgeTtsAdapter::build_ws_url("g", "c", "http://localhost:1234"); + assert!(url.starts_with("ws://localhost:1234")); + } + + #[test] + fn build_voices_url_contains_all_params() { + let url = EdgeTtsAdapter::build_voices_url("GEC123", "https://speech.platform.bing.com"); + assert!(url.starts_with("https://speech.platform.bing.com")); + assert!(url.contains(VOICES_PATH)); + assert!(url.contains("TrustedClientToken=")); + assert!(url.contains("Sec-MS-GEC=GEC123")); + } + + #[test] + fn build_ws_request_sets_origin_and_user_agent() { + let req = EdgeTtsAdapter::build_ws_request("wss://example.com/path").unwrap(); + assert_eq!( + req.headers().get("origin").unwrap(), + "https://speech.platform.bing.com" + ); + assert!(req.headers().get("user-agent").is_some()); + } + + // ============ SSML / 消息构造 ============ + + #[test] + fn build_ssml_contains_voice_and_text() { + let ssml = + EdgeTtsAdapter::build_ssml("hello", "zh-CN-XiaoxiaoNeural", "+0%", "+0Hz", "+0%"); + assert!(ssml.contains("zh-CN-XiaoxiaoNeural")); + assert!(ssml.contains("hello")); + assert!(ssml.contains("rate='+0%'")); + assert!(ssml.contains("pitch='+0Hz'")); + assert!(ssml.contains("volume='+0%'")); + assert!(ssml.contains("xml:lang='zh-CN'")); + } + + #[test] + fn build_ssml_escapes_xml_special_chars() { + let ssml = + EdgeTtsAdapter::build_ssml("ad", "zh-CN-XiaoxiaoNeural", "+0%", "+0Hz", "+0%"); + assert!(ssml.contains("<")); + assert!(ssml.contains(">")); + assert!(ssml.contains("&")); + assert!(!ssml.contains("ad"), "a<b&c>d"); + assert_eq!(escape_xml("plain"), "plain"); + } + + #[test] + fn generate_connection_id_is_32_hex() { + let id = generate_connection_id(); + assert_eq!(id.len(), 32); + assert!(id.chars().all(|c| c.is_ascii_hexdigit())); + } + + // ============ list_voices(HTTP,mockito) ============ + + #[tokio::test] + async fn list_voices_fetches_and_parses() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"ShortName": "zh-CN-XiaoxiaoNeural", "Name": "Microsoft Xiaoxiao", "Locale": "zh-CN", "Gender": "Female", "FriendlyName": "Xiaoxiao"}, + {"ShortName": "en-US-JennyNeural", "Name": "Microsoft Jenny", "Locale": "en-US", "Gender": "Female"} + ]); + server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let voices = adapter.list_voices(None).await.unwrap(); + assert_eq!(voices.len(), 2); + assert_eq!( + voices[0].short_name.as_deref(), + Some("zh-CN-XiaoxiaoNeural") + ); + assert_eq!(voices[0].gender.as_deref(), Some("Female")); + assert_eq!(voices[0].locale.as_deref(), Some("zh-CN")); + assert_eq!(voices[1].name.as_deref(), Some("Microsoft Jenny")); // 无 FriendlyName 回退 Name + } + + #[tokio::test] + async fn list_voices_filters_by_language() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"ShortName": "zh-CN-XiaoxiaoNeural", "Name": "x", "Locale": "zh-CN", "Gender": "Female"}, + {"ShortName": "zh-CN-YunxiNeural", "Name": "y", "Locale": "zh-CN", "Gender": "Male"}, + {"ShortName": "en-US-JennyNeural", "Name": "j", "Locale": "en-US", "Gender": "Female"} + ]); + server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let zh = adapter.list_voices(Some("zh-CN")).await.unwrap(); + assert_eq!(zh.len(), 2); + let en = adapter.list_voices(Some("en-US")).await.unwrap(); + assert_eq!(en.len(), 1); + let none = adapter.list_voices(Some("fr-FR")).await.unwrap(); + assert!(none.is_empty()); + } + + #[tokio::test] + async fn list_voices_caches_second_call() { + let mut server = mockito::Server::new_async().await; + let m = server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body( + serde_json::json!([{"ShortName":"v1","Name":"n","Locale":"zh-CN","Gender":"Female"}]) + .to_string(), + ) + .expect(1) // 第二次走缓存,不命中 HTTP + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let _ = adapter.list_voices(None).await.unwrap(); + let _ = adapter.list_voices(None).await.unwrap(); // 走缓存 + m.assert_async().await; + } + + #[tokio::test] + async fn list_voices_error_500_returns_service_unavailable() { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(503) + .with_body("unavailable") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ServiceUnavailable { .. })); + } + + // ============ recommend_voices ============ + + #[tokio::test] + async fn recommend_voices_filters_by_gender_and_limit() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"ShortName": "zh-CN-XiaoxiaoNeural", "Name": "x", "Locale": "zh-CN", "Gender": "Female"}, + {"ShortName": "zh-CN-YunxiNeural", "Name": "y", "Locale": "zh-CN", "Gender": "Male"}, + {"ShortName": "zh-CN-XiaoyiNeural", "Name": "z", "Locale": "zh-CN", "Gender": "Female"} + ]); + server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let female = adapter + .recommend_voices(Some("zh-CN"), Some("Female"), 10) + .await + .unwrap(); + assert_eq!(female.len(), 2); + let limited = adapter + .recommend_voices(Some("zh-CN"), None, 1) + .await + .unwrap(); + assert_eq!(limited.len(), 1); + } + + // ============ check_voice_available ============ + + #[tokio::test] + async fn check_voice_available_true_for_existing() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"ShortName": "zh-CN-XiaoxiaoNeural", "Name": "x", "Locale": "zh-CN", "Gender": "Female"} + ]); + server + .mock("GET", VOICES_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + assert!(adapter.check_voice_available("zh-CN-XiaoxiaoNeural").await); + assert!(!adapter.check_voice_available("zh-CN-Nonexistent").await); + } + + // ============ 音色降级(speech_with_fallback) ============ + + #[tokio::test] + async fn fallback_tries_next_on_voice_not_available() { + let count = Arc::new(AtomicU32::new(0)); + let count2 = count.clone(); + let single = move |v: &str| { + let v = v.to_string(); + let count3 = count2.clone(); + async move { + count3.fetch_add(1, Ordering::SeqCst); + if v == "v3" { + Ok(SpeechResult { + audio_data: Some(vec![1]), + ..Default::default() + }) + } else { + Err(AibridgeError::voice_not_available(&v)) + } + } + }; + let candidates = vec!["v1".to_string(), "v2".to_string(), "v3".to_string()]; + let result = EdgeTtsAdapter::speech_with_fallback(candidates, single).await; + assert!(result.is_ok()); + assert_eq!(count.load(Ordering::SeqCst), 3); + } + + #[tokio::test] + async fn fallback_returns_last_error_when_all_fail() { + let single = |v: &str| { + let v = v.to_string(); + async move { Err(AibridgeError::voice_not_available(&v)) } + }; + let candidates = vec!["v1".to_string(), "v2".to_string()]; + let result = EdgeTtsAdapter::speech_with_fallback(candidates, single).await; + assert!(matches!( + result, + Err(AibridgeError::VoiceNotAvailable { .. }) + )); + } + + #[tokio::test] + async fn fallback_single_voice_raises_directly() { + let single = |v: &str| { + let v = v.to_string(); + async move { Err(AibridgeError::voice_not_available(&v)) } + }; + let candidates = vec!["v1".to_string()]; + let result = EdgeTtsAdapter::speech_with_fallback(candidates, single).await; + assert!(matches!( + result, + Err(AibridgeError::VoiceNotAvailable { .. }) + )); + } + + #[tokio::test] + async fn fallback_non_fallback_error_raises_immediately() { + let count = Arc::new(AtomicU32::new(0)); + let count2 = count.clone(); + let single = move |_v: &str| { + let count3 = count2.clone(); + async move { + count3.fetch_add(1, Ordering::SeqCst); + Err(AibridgeError::authentication("auth fail")) // 非降级错误 + } + }; + let candidates = vec!["v1".to_string(), "v2".to_string()]; + let result = EdgeTtsAdapter::speech_with_fallback(candidates, single).await; + assert!(matches!(result, Err(AibridgeError::Authentication { .. }))); + assert_eq!(count.load(Ordering::SeqCst), 1); // 只调一次 + } + + #[tokio::test] + async fn fallback_first_success_skips_rest() { + let count = Arc::new(AtomicU32::new(0)); + let count2 = count.clone(); + let single = move |_v: &str| { + let count3 = count2.clone(); + async move { + count3.fetch_add(1, Ordering::SeqCst); + Ok(SpeechResult { + audio_data: Some(vec![1]), + ..Default::default() + }) + } + }; + let candidates = vec!["v1".to_string(), "v2".to_string(), "v3".to_string()]; + let result = EdgeTtsAdapter::speech_with_fallback(candidates, single).await; + assert!(result.is_ok()); + assert_eq!(count.load(Ordering::SeqCst), 1); // 第一个成功,不试后续 + } +} From 03776d26173c28c43d6d58211559a768c9faa9dd Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:29:54 +0800 Subject: [PATCH 41/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2b+2c=20=E6=94=B6=E5=B0=BE=20=E6=B3=A8=E5=86=8C=20kling=20+=20e?= =?UTF-8?q?dge-tts=20=E5=88=B0=E5=B7=A5=E5=8E=82=EF=BC=88edge-tts=20?= =?UTF-8?q?=E5=85=8D=E8=AE=A4=E8=AF=81=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 48 ++++++++++++++++++--- crates/aibridge-core/src/adapters/mod.rs | 6 +++ crates/aibridge-core/src/client.rs | 13 +++++- 3 files changed, 59 insertions(+), 8 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index c902c47..d4734c9 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -21,8 +21,10 @@ use crate::adapters::chinese::{ DoubaoAdapter, ErnieAdapter, KimiAdapter, MiniMaxAdapter, QwenAdapter, ZhipuAdapter, }; use crate::adapters::echo::EchoAdapter; +use crate::adapters::edge_tts::EdgeTtsAdapter; use crate::adapters::emerging_models::{IdeogramAdapter, LlamaAdapter, LumaAdapter}; use crate::adapters::gemini::GeminiAdapter; +use crate::adapters::kling::KlingAdapter; use crate::adapters::more_models::{ CohereAdapter, DeepSeekAdapter, MistralAdapter, PerplexityAdapter, StepFunAdapter, }; @@ -77,9 +79,11 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ // 阶段 2b 第二批独立协议(已实现): "runway", "pika", - // 阶段 2b/2c 待实现: + // 阶段 2b 第三批独立协议(已实现): "kling", + // 阶段 2c 音频(已实现): "edge-tts", + // 阶段 2c 待实现: "elevenlabs", "cartesia", "deepgram", @@ -140,10 +144,15 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { // 阶段 2b 第二批独立协议:别名对齐 Python agn/adapters/{runway,pika}.py 末尾 register 调用(均无别名) "runway" => Ok(Box::new(RunwayAdapter::new(config)?)), "pika" => Ok(Box::new(PikaAdapter::new(config)?)), - // 阶段 2 适配器占位 - "kling" | "edge-tts" | "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { + // 阶段 2b 第三批独立协议:别名对齐 Python agn/adapters/kling.py 末尾 register 调用(无别名) + "kling" => Ok(Box::new(KlingAdapter::new(config)?)), + // 阶段 2c 音频:别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用 + // (edge-tts / edge_tts / edge 三个别名均指向 EdgeTTSAdapter,免认证) + "edge-tts" | "edge_tts" | "edge" => Ok(Box::new(EdgeTtsAdapter::new(config)?)), + // 阶段 2c 适配器占位 + "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2 待实现)"), + provider: format!("{provider}(阶段 2c 待实现)"), }) } // 未知 provider @@ -443,6 +452,32 @@ mod tests { assert_eq!(adapter.provider_type(), "pika"); } + #[test] + fn create_kling_returns_adapter() { + // 阶段 2b:KlingAdapter 自带 DEFAULT_KLING_BASE_URL 兜底,仅需 api_key + let adapter = create_adapter(config_for("kling")).expect("工厂应能创建 kling 适配器"); + assert_eq!(adapter.provider_type(), "kling"); + } + + #[test] + fn create_edge_tts_returns_adapter() { + // 阶段 2c:EdgeTtsAdapter 免认证,无需 api_key(用 default opts 验证免认证构造) + let config = ProviderConfig::from_options("edge-tts", ClientOptions::default()); + let adapter = create_adapter(config).expect("工厂应能创建 edge-tts 适配器"); + assert_eq!(adapter.provider_type(), "edge-tts"); + } + + #[test] + fn create_edge_tts_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用: + // edge_tts / edge -> edge-tts(均指向 EdgeTTSAdapter,免认证) + let edge_tts = + create_adapter(config_for("edge_tts")).expect("别名 edge_tts 应映射到 edge-tts"); + assert_eq!(edge_tts.provider_type(), "edge-tts"); + let edge = create_adapter(config_for("edge")).expect("别名 edge 应映射到 edge-tts"); + assert_eq!(edge.provider_type(), "edge-tts"); + } + #[test] fn create_additional_models_aliases_map_to_main_provider_type() { // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: @@ -470,8 +505,8 @@ mod tests { #[test] fn create_phase2_adapter_returns_phase2_message() { - // kling 仍为阶段 2 占位(未实现),返 ProviderNotFound - let result = create_adapter(config_for("kling")); + // elevenlabs 仍为阶段 2c 占位(未实现),返 ProviderNotFound + let result = create_adapter(config_for("elevenlabs")); if let Err(AibridgeError::ProviderNotFound { provider }) = result { assert!(provider.contains("阶段 2")); } else { @@ -485,6 +520,7 @@ mod tests { assert!(is_known_provider("openai")); assert!(is_known_provider("edge-tts")); assert!(is_known_provider("assemblyai")); + assert!(is_known_provider("kling")); // 阶段 2a 已实现 provider 应被识别 assert!(is_known_provider("azure")); assert!(is_known_provider("siliconflow")); diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 0800bea..5a3d71f 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -57,3 +57,9 @@ pub mod runway; /// Pika 适配器:阶段 2b 独立协议,视频生成(文生视频/图生视频/任务轮询) pub mod pika; + +/// Kling 适配器:阶段 2b 独立协议,可灵视频生成(文生视频/图生视频/任务轮询) +pub mod kling; + +/// Edge TTS 适配器:阶段 2c 音频,免费文字转语音(免认证,WebSocket 协议) +pub mod edge_tts; diff --git a/crates/aibridge-core/src/client.rs b/crates/aibridge-core/src/client.rs index 717320f..a684295 100644 --- a/crates/aibridge-core/src/client.rs +++ b/crates/aibridge-core/src/client.rs @@ -69,9 +69,9 @@ impl Client { /// 判断 provider 是否为免认证(不需要 API Key) /// /// 阶段 0.6:`echo` 为 mock 适配器,免认证便于五语言管线验证。 - /// 阶段 2c 起将补充 edge-tts 等真实免认证 provider。 + /// 阶段 2c:`edge-tts` 为真实免认证 provider(微软免费 TTS,含别名 edge_tts / edge)。 fn is_free_provider(provider: &str) -> bool { - matches!(provider, "echo") + matches!(provider, "echo" | "edge-tts" | "edge_tts" | "edge") } /// 启动客户端(初始化适配器) @@ -207,6 +207,15 @@ mod tests { assert!(result.is_ok(), "echo 应免认证创建客户端"); } + #[test] + fn new_edge_tts_without_api_key_succeeds() { + // edge-tts 免认证:无 api_key 也能创建客户端(含别名 edge_tts / edge) + let result = Client::new("edge-tts", ClientOptions::default()); + assert!(result.is_ok(), "edge-tts 应免认证创建客户端"); + let result_alias = Client::new("edge", ClientOptions::default()); + assert!(result_alias.is_ok(), "edge 别名应免认证创建客户端"); + } + #[tokio::test] async fn echo_client_chat_roundtrip() { // 端到端验证:Client → EchoAdapter → chat 回显 From 0b7c205379c2931e34ad8e89cb50a5be41a77923 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:45:49 +0800 Subject: [PATCH 42/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20elevenlabs=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit ElevenLabs TTS 适配器(独立协议,xi-api-key 认证)。 - speech: POST /text-to-speech/{voice_id},二进制音频返回 - list_voices: GET /voices,按语言过滤;recommend_voices 复用默认实现按性别+limit 过滤 - 错误映射: 401→Authentication / 429→RateLimit / 400→Validation / 404→VoiceNotAvailable / 5xx→Api - 内置 46 个音色别名表(Rachel→voice_id),未知名称原样透传 - 64 个单测,覆盖正常/错误路径(401/429/400/404/500)/不支持能力/纯函数 --- .../aibridge-core/src/adapters/elevenlabs.rs | 1467 +++++++++++++++++ 1 file changed, 1467 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/elevenlabs.rs diff --git a/crates/aibridge-core/src/adapters/elevenlabs.rs b/crates/aibridge-core/src/adapters/elevenlabs.rs new file mode 100644 index 0000000..1b902bd --- /dev/null +++ b/crates/aibridge-core/src/adapters/elevenlabs.rs @@ -0,0 +1,1467 @@ +//! ElevenLabs 适配器(高质量多语言 TTS) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/audio_adapters.py` 的 `ElevenLabsAdapter`。 +//! +//! ElevenLabs 是全球流行的高质量语音合成平台,支持超逼真音色、多语言、声音克隆。 +//! +//! ## 协议(独立协议,非 OpenAI 兼容) +//! +//! - Base URL: `https://api.elevenlabs.io/v1` +//! - 认证: `xi-api-key` 请求头(非 Bearer) +//! - speech: `POST /text-to-speech/{voice_id}` +//! - 请求体: `{ "text", "model_id", "voice_settings"?: { stability, similarity_boost, style, use_speaker_boost } }` +//! - 查询参数: `output_format`(非 mp3 时携带,如 wav/ogg/opus/ulaw/pcm) +//! - 响应: 二进制音频流(mp3 默认) +//! - list_voices: `GET /voices`,返回 `{ "voices": [{ voice_id, name, labels, category }] }` +//! - list_models: 无标准 /models 拉取需求,保留硬编码 TTS 模型列表 +//! +//! ## 错误映射 +//! +//! - 401/403 → Authentication(xi-api-key 无效) +//! - 429 → RateLimit(限流或配额耗尽) +//! - 400/422 → Validation(请求参数错误) +//! - 404 → VoiceNotAvailable(voice_id 或模型不存在) +//! - 5xx → Api(服务端错误) +//! - 其余 4xx → Api +//! +//! ## 特性(v1.3.3 保留) +//! +//! - `requires_api_key = true`(需 API Key) +//! - 内置音色别名表("Rachel" → voice_id),未知名称视为直接传 voice_id +//! - 收到候选音色列表时取第一个(不实现多音色降级,与 Python v1 一致) + +use std::collections::HashMap; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::Value; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::audio::{SpeechRequest, SpeechResult}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; + +// ==================== 常量 ==================== + +/// Provider 类型标识 +const PROVIDER_TYPE: &str = "elevenlabs"; +/// Provider 显示名称 +const PROVIDER_NAME: &str = "ElevenLabs"; +/// 默认 API 基地址 +const DEFAULT_API_BASE: &str = "https://api.elevenlabs.io/v1"; +/// 默认模型(多语言 v2,ElevenLabs 推荐) +const DEFAULT_MODEL: &str = "eleven_multilingual_v2"; +/// 默认音色 voice_id(Rachel,ElevenLabs 最常用女声) +const DEFAULT_VOICE_ID: &str = "21m00Tcm4TlvDq8ikWAM"; +/// TTS 合成端点路径前缀 +const TTS_PATH: &str = "/text-to-speech/"; +/// 音色列表端点路径 +const VOICES_PATH: &str = "/voices"; + +// ==================== ElevenLabs 原始音色反序列化结构 ==================== + +/// ElevenLabs /voices 返回的单个音色项(原始 JSON 结构) +/// +/// 字段名与 ElevenLabs 服务端返回的 JSON 一致(snake_case)。 +#[derive(Debug, Deserialize)] +struct ElevenVoiceRaw { + /// 音色 ID(如 "21m00Tcm4TlvDq8ikWAM") + voice_id: String, + /// 音色显示名(如 "Rachel") + name: String, + /// 标签集合(含 gender / language / accent 等元数据) + #[serde(default)] + labels: HashMap, + /// 音色类别(如 "premade" / "cloned") + #[serde(default)] + category: Option, +} + +/// ElevenLabs /voices 响应外壳 +#[derive(Debug, Deserialize)] +struct ElevenVoicesResponse { + /// 音色列表 + voices: Vec, +} + +// ==================== ElevenLabsAdapter ==================== + +/// ElevenLabs 适配器 +/// +/// 持有 HTTP 客户端与 API Key。所有请求按需发起(无长连接)。 +pub struct ElevenLabsAdapter { + /// Provider 配置 + #[allow(dead_code)] + config: ProviderConfig, + /// HTTP 客户端 + http: HttpClient, + /// API Key(xi-api-key 请求头用) + api_key: Option, + /// API 基地址(默认 `https://api.elevenlabs.io/v1`,可由 config.base_url 覆盖) + api_base: String, +} + +impl ElevenLabsAdapter { + /// 创建 ElevenLabs 适配器 + /// + /// `config.base_url` 可覆盖 API 基地址(主要用于测试指向 mock server), + /// 为空时用 `DEFAULT_API_BASE`。API Key 从 `config.api_key` 取。 + pub fn new(config: ProviderConfig) -> Result { + let api_base = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_BASE.to_string()); + let api_key = config.api_key.clone(); + let opts = ClientOptions::builder().timeout(config.timeout).build(); + let http = HttpClient::new(&opts)?; + Ok(Self { + config, + http, + api_key, + api_base, + }) + } + + // ==================== 纯函数(协议构造/解析,可单测) ==================== + + /// 解析音色:别名(如 "Rachel")→ voice_id,未知则原样返回(视为直接传了 voice_id) + /// + /// 对应 Python v1 `ElevenLabsAdapter._get_voice_id`。 + /// - 空值 → 默认音色 Rachel 的 voice_id + /// - 已知别名(大小写不敏感)→ 对应 voice_id + /// - 其余 → 原样返回 + fn resolve_voice_id(voice: &str) -> String { + if voice.is_empty() { + return DEFAULT_VOICE_ID.to_string(); + } + if let Some(id) = lookup_default_voice(voice) { + return id.to_string(); + } + voice.to_string() + } + + /// 构造 speech 请求体 + /// + /// 对应 Python v1 `ElevenLabsAdapter.speech` 的 payload 构造。 + /// 从 `extra` 提取 ElevenLabs 特有参数(stability/similarity_boost/style/use_speaker_boost) + /// 组装到 `voice_settings`。 + fn build_speech_payload(input: &str, model: &str, extra: &HashMap) -> Value { + let mut payload = serde_json::json!({ + "text": input, + "model_id": model, + }); + let mut voice_settings = serde_json::Map::new(); + for key in [ + "stability", + "similarity_boost", + "style", + "use_speaker_boost", + ] { + if let Some(v) = extra.get(key) { + voice_settings.insert(key.to_string(), v.clone()); + } + } + if !voice_settings.is_empty() { + payload["voice_settings"] = Value::Object(voice_settings); + } + payload + } + + /// 根据输出格式返回对应 Content-Type + /// + /// 对应 Python v1 `ElevenLabsAdapter.speech` 的 `content_type_map`。 + fn content_type_for_format(fmt: &str) -> &'static str { + match fmt.to_lowercase().as_str() { + "mp3" => "audio/mpeg", + "wav" => "audio/wav", + "ogg" => "audio/ogg", + "opus" => "audio/opus", + "ulaw" => "audio/basic", + "pcm" => "audio/pcm", + _ => "audio/mpeg", + } + } + + /// 把 ElevenLabs 原始音色转为统一 `VoiceInfo` + /// + /// - voice_id → voice_id + short_name(ElevenLabs 用 voice_id 作唯一标识) + /// - name → name + /// - labels.gender → gender(标准化为 "Female"/"Male") + /// - labels.language → locale(标准化为语言代码如 "en"/"zh") + /// - labels 其余字段 + category → extra + fn convert_voice(raw: ElevenVoiceRaw) -> VoiceInfo { + let gender = raw.labels.get("gender").map(|g| capitalize_gender(g)); + let locale = raw + .labels + .get("language") + .map(|l| normalize_language_code(l)); + let voice_id = raw.voice_id.clone(); + let mut builder = VoiceInfo::builder() + .voice_id(voice_id.clone()) + .short_name(voice_id) + .name(raw.name); + if let Some(g) = gender { + builder = builder.gender(g); + } + if let Some(l) = locale { + builder = builder.locale(l); + } + let mut info = builder.build(); + for (k, val) in raw.labels { + info.extra.insert(k, Value::String(val)); + } + if let Some(cat) = raw.category { + info.extra + .insert("category".to_string(), Value::String(cat)); + } + info + } + + /// 将 ElevenLabs HTTP 错误响应映射为 AibridgeError + /// + /// 映射规则见模块文档"错误映射"小节。 + fn map_elevenlabs_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 | 403 => AibridgeError::authentication(format!( + "ElevenLabs 认证失败(xi-api-key 无效): {body}" + )), + 429 => AibridgeError::rate_limit(format!("ElevenLabs 限流或配额耗尽: {body}")), + 400 | 422 => AibridgeError::validation(format!("ElevenLabs 请求参数错误: {body}")), + 404 => AibridgeError::voice_not_available(format!( + "ElevenLabs voice_id 或模型不存在: {body}" + )), + s if s >= 500 => AibridgeError::api(s, format!("ElevenLabs 服务错误 ({s}): {body}")), + s => AibridgeError::api(s, format!("ElevenLabs HTTP {s}: {body}")), + } + } + + /// 构造 speech 端点完整 URL(`{api_base}/text-to-speech/{voice_id}`) + fn build_tts_url(api_base: &str, voice_id: &str) -> String { + format!( + "{base}{path}{voice_id}", + base = api_base.trim_end_matches('/'), + path = TTS_PATH, + voice_id = voice_id, + ) + } + + /// 构造 list_voices 端点完整 URL(`{api_base}/voices`) + fn build_voices_url(api_base: &str) -> String { + format!( + "{base}{path}", + base = api_base.trim_end_matches('/'), + path = VOICES_PATH, + ) + } + + /// 获取 API Key,为空时返 Validation 错误 + fn require_api_key(&self) -> Result<&str> { + self.api_key + .as_deref() + .filter(|k| !k.trim().is_empty()) + .ok_or_else(|| AibridgeError::validation("ElevenLabs 需要 API key(xi-api-key)")) + } +} + +#[async_trait] +impl Adapter for ElevenLabsAdapter { + fn provider_type(&self) -> &str { + PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::AudioSpeech); + caps.insert(Capabilities::ListVoices); + caps + } + + /// 需 API Key 认证 + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new 时已建好,保持 no-op + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // 无长连接资源需释放 + Ok(()) + } + + /// 文字转语音(ElevenLabs TTS) + /// + /// 收到候选音色列表时取第一个(不实现多音色降级,与 Python v1 一致)。 + async fn speech(&self, req: SpeechRequest) -> Result { + let api_key = self.require_api_key()?; + + // 取主音色并解析为 voice_id + let voice_raw = req.voice.primary().unwrap_or_default(); + let voice_id = Self::resolve_voice_id(voice_raw); + + // 模型缺省回退到默认模型 + let model = if req.model.is_empty() { + DEFAULT_MODEL.to_string() + } else { + req.model.clone() + }; + + // 输出格式(默认 mp3) + let fmt = if req.response_format.is_empty() { + "mp3" + } else { + req.response_format.as_str() + }; + let content_type = Self::content_type_for_format(fmt); + + let payload = Self::build_speech_payload(&req.input, &model, &req.extra); + let url = Self::build_tts_url(&self.api_base, &voice_id); + + // 构造请求:xi-api-key + Accept + JSON body;非 mp3 时附加 output_format 查询参数 + let mut request = self + .http + .inner() + .post(&url) + .header("xi-api-key", api_key) + .header("Accept", content_type) + .json(&payload); + if fmt != "mp3" { + request = request.query(&[("output_format", fmt)]); + } + + let resp = request.send().await.map_err(map_reqwest_err)?; + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + return Err(Self::map_elevenlabs_error(status.as_u16(), &body)); + } + + let audio_bytes = resp.bytes().await.map_err(map_reqwest_err)?; + + Ok(SpeechResult { + audio_data: Some(audio_bytes.to_vec()), + audio_url: None, + audio_base64: None, + content_type: content_type.to_string(), + format: fmt.to_string(), + duration: None, + model: Some(model), + }) + } + + /// 列出可用音色(按 language 过滤) + /// + /// language 过滤对 labels.language 做标准化后前缀匹配(如 "en" 匹配 "english")。 + async fn list_voices(&self, language: Option<&str>) -> Result> { + let api_key = self.require_api_key()?; + let url = Self::build_voices_url(&self.api_base); + + let resp = self + .http + .inner() + .get(&url) + .header("xi-api-key", api_key) + .send() + .await + .map_err(map_reqwest_err)?; + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + return Err(Self::map_elevenlabs_error(status.as_u16(), &body)); + } + + let bytes = resp.bytes().await.map_err(map_reqwest_err)?; + let resp_json: ElevenVoicesResponse = serde_json::from_slice(&bytes).map_err(|e| { + AibridgeError::validation(format!("ElevenLabs /voices 响应解析失败: {e}")) + })?; + let voices: Vec = resp_json + .voices + .into_iter() + .map(Self::convert_voice) + .collect(); + + match language { + Some(lang) if !lang.is_empty() => Ok(voices + .into_iter() + .filter(|v| voice_matches_language(lang, v)) + .collect()), + _ => Ok(voices), + } + } + + /// 列出 ElevenLabs TTS 模型(硬编码列表,无 /models 拉取需求) + async fn list_models(&self, filter: Option) -> Result> { + let models = elevenlabs_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 辅助函数 ==================== + +/// 查询内置音色别名 → voice_id +/// +/// 对应 Python v1 `ElevenLabsAdapter.DEFAULT_VOICES`。返回 None 表示非已知别名。 +fn lookup_default_voice(voice: &str) -> Option<&'static str> { + let lower = voice.to_lowercase(); + let id = match lower.as_str() { + "rachel" => "21m00Tcm4TlvDq8ikWAM", + "drew" => "29vD33N1CtxCmqQRPOHJ", + "clyde" => "2EiwWnXFnvU5JabPnv8n", + "paul" => "5Q0t7uMcjvnagumLfvZi", + "domi" => "AZnzlk1XvdvUeBnXmlld", + "dave" => "CYw3kZ02Hs0563khs1Fj", + "fin" => "D38z5RcWu1voky8WS1ja", + "sarah" => "EXAVITQu4vr4xnSDxMaL", + "antoni" => "ErXwobaYiN019PkySvjV", + "thomas" => "GBv7mTt0atIp3Br8iCZE", + "charlie" => "IKne3meq5aSn9XLyUdCD", + "george" => "JBFqnCBsd32t6Ie9FZ2Q", + "emily" => "LcfcDJNUP1GQjkzn1xUU", + "elli" => "MF3mGyEYCl7XYWbV9V6O", + "callum" => "N2lVS1w4EtoT3dr4eOWO", + "patrick" => "ODq5zmih8GrVes37DekR", + "harry" => "SOYHLrjzK2X1ezoPC6cr", + "liam" => "TX3LPaxmHKxFdv7VOQHJ", + "dorothy" => "ThT5KcBeYPX3keUQqHPh", + "josh" => "TxGEqnHWrfWFTfGW9XjX", + "arnold" => "VR6AewLTigWG4xSOukaG", + "charlotte" => "XB0fDUnXU5powFXDhCwa", + "alice" => "Xb7hH8MSUJpSbSDYk0k2", + "matilda" => "XrExE9yKIg1WjnnlVkGX", + "matthew" => "Yko7PKHZNXotIFUBG7I9", + "james" => "ZQe5CZNOzWyzPSCn5a3c", + "joseph" => "Zlb1dXrM653N07WRPnSh", + "jeremy" => "bVMeCyTHy58xNoL34h3p", + "michael" => "flq6f7yk4E4fJM5XTYuZ", + "ethan" => "g5CIjZEefAph4nQFvHAz", + "chris" => "iP95p4xoKVk53GoZ742B", + "gigi" => "jBpfuIE2acCO8z3wKNLl", + "freya" => "jsCqWAovK2LkecY7zXl4", + "brian" => "nPczCjzI2devxB14oouP", + "grace" => "oWAxZDx7w5VEj9dCyTzz", + "daniel" => "onwK4e9ZLuTAKqWW03F9", + "lily" => "pFZP5JQG7iQjIQuC4Bku", + "serena" => "pMsXgVXv3BLzUgSXRplE", + "adam" => "pNInz6obpgDQGcFmaJgB", + "nicole" => "piTKgcLEGmPE4e6mEKli", + "bill" => "pqHfZKP75CvOlQylNhV4", + "jessie" => "t0jbNlBVZ17f02VDIeMI", + "sam" => "yoZ06aMxZJJ28mfd3POQ", + "glinda" => "z9fAnlkpzviVjnFOo0Tc", + "giovanni" => "zcAOhNBS3c14rBihAFp1", + "mimi" => "zrHiDhphv9ZnVXBqCLjz", + _ => return None, + }; + Some(id) +} + +/// 性别字符串标准化(female/f → Female,male/m → Male) +fn capitalize_gender(g: &str) -> String { + match g.to_lowercase().as_str() { + "female" | "f" | "woman" => "Female".to_string(), + "male" | "m" | "man" => "Male".to_string(), + _ => g.to_string(), + } +} + +/// 语言名称/代码标准化为 2 字母代码(便于前缀匹配) +/// +/// ElevenLabs labels.language 可能是自然语言名("english"/"chinese")或代码("en"/"zh-CN"), +/// 统一映射到 2 字母代码做过滤匹配。 +fn normalize_language_code(s: &str) -> String { + let lower = s.to_lowercase(); + let code = match lower.as_str() { + "english" | "en" | "en-us" | "en-gb" | "en-au" | "en-in" => "en", + "chinese" | "zh" | "zh-cn" | "zh-tw" | "mandarin" | "cantonese" => "zh", + "japanese" | "ja" | "ja-jp" => "ja", + "korean" | "ko" | "ko-kr" => "ko", + "spanish" | "es" | "es-es" | "es-mx" => "es", + "french" | "fr" | "fr-fr" | "fr-ca" => "fr", + "german" | "de" | "de-de" => "de", + "italian" | "it" | "it-it" => "it", + "portuguese" | "pt" | "pt-br" | "pt-pt" => "pt", + "russian" | "ru" | "ru-ru" => "ru", + "hindi" | "hi" => "hi", + "arabic" | "ar" => "ar", + "dutch" | "nl" => "nl", + "polish" | "pl" => "pl", + "turkish" | "tr" => "tr", + "indonesian" | "id" => "id", + "vietnamese" | "vi" => "vi", + "thai" | "th" => "th", + _ => &lower, + }; + code.to_string() +} + +/// 判断音色是否匹配指定语言(标准化后前缀互匹配) +/// +/// 用于 list_voices 的 language 过滤。filter 与 voice 的 locale 双向前缀匹配, +/// 兼容 "en" ↔ "english"(均标准化为 "en")。 +fn voice_matches_language(filter: &str, voice: &VoiceInfo) -> bool { + let f = normalize_language_code(filter); + if f.is_empty() { + return true; + } + voice + .locale + .as_deref() + .map(|l| { + let norm = normalize_language_code(l); + !norm.is_empty() && (norm.starts_with(&f) || f.starts_with(&norm)) + }) + .unwrap_or(false) +} + +/// 将 reqwest::Error 映射为 AibridgeError(超时 → Timeout,其余 → Network) +fn map_reqwest_err(e: reqwest::Error) -> AibridgeError { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } +} + +/// ElevenLabs TTS 模型硬编码列表 +fn elevenlabs_models() -> Vec { + vec![ + ModelInfo { + id: "eleven_multilingual_v2".into(), + name: "Eleven Multilingual v2".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: true, + description: Some("ElevenLabs 推荐的多语言 TTS 模型,支持 29 种语言,质量最高".into()), + created: None, + }, + ModelInfo { + id: "eleven_turbo_v2".into(), + name: "Eleven Turbo v2".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: true, + description: Some("低延迟 TTS 模型,兼顾质量与速度,适合实时场景".into()), + created: None, + }, + ModelInfo { + id: "eleven_flash_v2_5".into(), + name: "Eleven Flash v2.5".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: true, + description: Some("最快 TTS 模型,极低延迟,适合高并发场景".into()), + created: None, + }, + ModelInfo { + id: "eleven_multilingual_v1".into(), + name: "Eleven Multilingual v1".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: true, + description: Some("多语言 TTS 模型 v1(已弃用,建议用 v2)".into()), + created: None, + }, + ModelInfo { + id: "eleven_monolingual_v1".into(), + name: "Eleven Monolingual v1".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: true, + description: Some("单语种(英文)TTS 模型 v1".into()), + created: None, + }, + ] +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::model::audio::TranscribeRequest; + use crate::model::chat::ChatRequest; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + + /// 构造测试用适配器(base_url 指向 mockito server,api_key 可选) + fn make_adapter(base_url: Option, api_key: Option<&str>) -> ElevenLabsAdapter { + let mut opts = ClientOptions::builder(); + if let Some(u) = base_url { + opts = opts.base_url(u); + } + if let Some(k) = api_key { + opts = opts.api_key(k); + } + let config = ProviderConfig::from_options(PROVIDER_TYPE, opts.build()); + ElevenLabsAdapter::new(config).expect("构造 ElevenLabsAdapter 失败") + } + + // ============ 基本属性 ============ + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter(None, Some("test-key")); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_speech_and_list_voices() { + let adapter = make_adapter(None, Some("test-key")); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::AudioSpeech)); + assert!(caps.contains(&Capabilities::ListVoices)); + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::AudioTranscribe)); + } + + #[test] + fn provider_type_and_name() { + let adapter = make_adapter(None, Some("test-key")); + assert_eq!(adapter.provider_type(), "elevenlabs"); + assert_eq!(adapter.provider_name(), "ElevenLabs"); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter(None, Some("test-key")); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_returns_elevenlabs_models() { + let adapter = make_adapter(None, Some("test-key")); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 5); + assert_eq!(models[0].id, "eleven_multilingual_v2"); + assert_eq!(models[0].model_type, ModelType::Audio); + assert_eq!(models[0].provider, "elevenlabs"); + } + + #[tokio::test] + async fn list_models_filter_by_audio() { + let adapter = make_adapter(None, Some("test-key")); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 5); + let chat = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat.is_empty()); + } + + // ============ 不支持能力(默认实现) ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = ChatRequest::builder("m", vec![]).build(); + assert!(matches!( + adapter.chat(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn image_generate_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = ImageRequest::builder("m", "p").build(); + assert!(matches!( + adapter.image_generate(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = VideoRequest::builder("m", "p").build(); + assert!(matches!( + adapter.video_create(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn embed_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + assert!(matches!( + adapter.embed(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn transcribe_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = TranscribeRequest::builder("m", FileInput::path("/tmp/a.mp3")).build(); + assert!(matches!( + adapter.transcribe(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + // ============ resolve_voice_id ============ + + #[test] + fn resolve_voice_id_empty_returns_default() { + assert_eq!(ElevenLabsAdapter::resolve_voice_id(""), DEFAULT_VOICE_ID); + } + + #[test] + fn resolve_voice_id_known_alias() { + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("Rachel"), + "21m00Tcm4TlvDq8ikWAM" + ); + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("Antoni"), + "ErXwobaYiN019PkySvjV" + ); + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("Adam"), + "pNInz6obpgDQGcFmaJgB" + ); + } + + #[test] + fn resolve_voice_id_case_insensitive() { + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("rachel"), + "21m00Tcm4TlvDq8ikWAM" + ); + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("RACHEL"), + "21m00Tcm4TlvDq8ikWAM" + ); + } + + #[test] + fn resolve_voice_id_unknown_passes_through() { + // 未知别名视为直接传了 voice_id,原样返回 + assert_eq!( + ElevenLabsAdapter::resolve_voice_id("abc123XYZ"), + "abc123XYZ" + ); + } + + // ============ content_type_for_format ============ + + #[test] + fn content_type_for_format_mp3() { + assert_eq!( + ElevenLabsAdapter::content_type_for_format("mp3"), + "audio/mpeg" + ); + assert_eq!( + ElevenLabsAdapter::content_type_for_format("MP3"), + "audio/mpeg" + ); + } + + #[test] + fn content_type_for_format_various() { + assert_eq!( + ElevenLabsAdapter::content_type_for_format("wav"), + "audio/wav" + ); + assert_eq!( + ElevenLabsAdapter::content_type_for_format("ogg"), + "audio/ogg" + ); + assert_eq!( + ElevenLabsAdapter::content_type_for_format("opus"), + "audio/opus" + ); + assert_eq!( + ElevenLabsAdapter::content_type_for_format("ulaw"), + "audio/basic" + ); + assert_eq!( + ElevenLabsAdapter::content_type_for_format("pcm"), + "audio/pcm" + ); + } + + #[test] + fn content_type_for_format_unknown_defaults_mp3() { + assert_eq!( + ElevenLabsAdapter::content_type_for_format("xyz"), + "audio/mpeg" + ); + } + + // ============ build_speech_payload ============ + + #[test] + fn build_speech_payload_minimal() { + let payload = ElevenLabsAdapter::build_speech_payload( + "hello", + "eleven_multilingual_v2", + &HashMap::new(), + ); + assert_eq!(payload["text"], "hello"); + assert_eq!(payload["model_id"], "eleven_multilingual_v2"); + assert!(payload.get("voice_settings").is_none()); + } + + #[test] + fn build_speech_payload_with_voice_settings() { + let mut extra = HashMap::new(); + extra.insert("stability".to_string(), serde_json::json!(0.5)); + extra.insert("similarity_boost".to_string(), serde_json::json!(0.75)); + extra.insert("style".to_string(), serde_json::json!(0.0)); + extra.insert("use_speaker_boost".to_string(), serde_json::json!(true)); + + let payload = ElevenLabsAdapter::build_speech_payload("hi", "eleven_turbo_v2", &extra); + let vs = payload.get("voice_settings").expect("应有 voice_settings"); + assert_eq!(vs["stability"], 0.5); + assert_eq!(vs["similarity_boost"], 0.75); + assert_eq!(vs["style"], 0.0); + assert_eq!(vs["use_speaker_boost"], true); + } + + #[test] + fn build_speech_payload_ignores_unrelated_extra() { + let mut extra = HashMap::new(); + extra.insert("stability".to_string(), serde_json::json!(0.3)); + extra.insert("unrelated_key".to_string(), serde_json::json!("ignored")); + + let payload = ElevenLabsAdapter::build_speech_payload("hi", "m", &extra); + let vs = payload.get("voice_settings").unwrap(); + assert_eq!(vs["stability"], 0.3); + assert!(vs.get("unrelated_key").is_none()); + } + + // ============ map_elevenlabs_error ============ + + #[test] + fn map_error_401_is_authentication() { + let err = ElevenLabsAdapter::map_elevenlabs_error(401, "invalid key"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_error_403_is_authentication() { + let err = ElevenLabsAdapter::map_elevenlabs_error(403, "forbidden"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_error_429_is_rate_limit() { + let err = ElevenLabsAdapter::map_elevenlabs_error(429, "slow down"); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn map_error_400_is_validation() { + let err = ElevenLabsAdapter::map_elevenlabs_error(400, "bad request"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_error_422_is_validation() { + let err = ElevenLabsAdapter::map_elevenlabs_error(422, "unprocessable"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_error_404_is_voice_not_available() { + let err = ElevenLabsAdapter::map_elevenlabs_error(404, "voice not found"); + assert!(matches!(err, AibridgeError::VoiceNotAvailable { .. })); + } + + #[test] + fn map_error_500_is_api() { + let err = ElevenLabsAdapter::map_elevenlabs_error(500, "server error"); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[test] + fn map_error_503_is_api() { + let err = ElevenLabsAdapter::map_elevenlabs_error(503, "unavailable"); + assert!(matches!(err, AibridgeError::Api { status: 503, .. })); + } + + #[test] + fn map_error_other_4xx_is_api() { + let err = ElevenLabsAdapter::map_elevenlabs_error(418, "teapot"); + assert!(matches!(err, AibridgeError::Api { status: 418, .. })); + } + + // ============ URL 构造 ============ + + #[test] + fn build_tts_url_contains_voice_id() { + let url = ElevenLabsAdapter::build_tts_url( + "https://api.elevenlabs.io/v1", + "21m00Tcm4TlvDq8ikWAM", + ); + assert_eq!( + url, + "https://api.elevenlabs.io/v1/text-to-speech/21m00Tcm4TlvDq8ikWAM" + ); + } + + #[test] + fn build_tts_url_strips_trailing_slash() { + let url = ElevenLabsAdapter::build_tts_url("https://api.elevenlabs.io/v1/", "vid"); + assert_eq!(url, "https://api.elevenlabs.io/v1/text-to-speech/vid"); + } + + #[test] + fn build_voices_url_is_correct() { + let url = ElevenLabsAdapter::build_voices_url("https://api.elevenlabs.io/v1"); + assert_eq!(url, "https://api.elevenlabs.io/v1/voices"); + } + + // ============ convert_voice ============ + + #[test] + fn convert_voice_maps_all_fields() { + let mut labels = HashMap::new(); + labels.insert("gender".to_string(), "female".to_string()); + labels.insert("language".to_string(), "english".to_string()); + labels.insert("accent".to_string(), "american".to_string()); + let raw = ElevenVoiceRaw { + voice_id: "21m00Tcm4TlvDq8ikWAM".into(), + name: "Rachel".into(), + labels, + category: Some("premade".into()), + }; + let v = ElevenLabsAdapter::convert_voice(raw); + assert_eq!(v.voice_id.as_deref(), Some("21m00Tcm4TlvDq8ikWAM")); + assert_eq!(v.short_name.as_deref(), Some("21m00Tcm4TlvDq8ikWAM")); + assert_eq!(v.name.as_deref(), Some("Rachel")); + assert_eq!(v.gender.as_deref(), Some("Female")); + assert_eq!(v.locale.as_deref(), Some("en")); + assert_eq!( + v.extra.get("accent").and_then(|x| x.as_str()), + Some("american") + ); + assert_eq!( + v.extra.get("category").and_then(|x| x.as_str()), + Some("premade") + ); + } + + #[test] + fn convert_voice_without_labels() { + let raw = ElevenVoiceRaw { + voice_id: "v1".into(), + name: "Voice".into(), + labels: HashMap::new(), + category: None, + }; + let v = ElevenLabsAdapter::convert_voice(raw); + assert_eq!(v.voice_id.as_deref(), Some("v1")); + assert_eq!(v.name.as_deref(), Some("Voice")); + assert!(v.gender.is_none()); + assert!(v.locale.is_none()); + } + + #[test] + fn convert_voice_male_gender() { + let mut labels = HashMap::new(); + labels.insert("gender".to_string(), "male".to_string()); + labels.insert("language".to_string(), "chinese".to_string()); + let raw = ElevenVoiceRaw { + voice_id: "v2".into(), + name: "MaleVoice".into(), + labels, + category: None, + }; + let v = ElevenLabsAdapter::convert_voice(raw); + assert_eq!(v.gender.as_deref(), Some("Male")); + assert_eq!(v.locale.as_deref(), Some("zh")); + } + + // ============ 辅助函数 ============ + + #[test] + fn capitalize_gender_variants() { + assert_eq!(capitalize_gender("female"), "Female"); + assert_eq!(capitalize_gender("FEMALE"), "Female"); + assert_eq!(capitalize_gender("male"), "Male"); + assert_eq!(capitalize_gender("m"), "Male"); + assert_eq!(capitalize_gender("unknown"), "unknown"); + } + + #[test] + fn normalize_language_code_known_names() { + assert_eq!(normalize_language_code("english"), "en"); + assert_eq!(normalize_language_code("chinese"), "zh"); + assert_eq!(normalize_language_code("japanese"), "ja"); + assert_eq!(normalize_language_code("korean"), "ko"); + assert_eq!(normalize_language_code("spanish"), "es"); + assert_eq!(normalize_language_code("french"), "fr"); + assert_eq!(normalize_language_code("german"), "de"); + } + + #[test] + fn normalize_language_code_codes() { + assert_eq!(normalize_language_code("en"), "en"); + assert_eq!(normalize_language_code("en-US"), "en"); + assert_eq!(normalize_language_code("zh-CN"), "zh"); + assert_eq!(normalize_language_code("ja-JP"), "ja"); + } + + #[test] + fn normalize_language_code_unknown_passthrough() { + assert_eq!(normalize_language_code("klingon"), "klingon"); + } + + #[test] + fn voice_matches_language_en_matches_english() { + let v = VoiceInfo::builder().locale("en").build(); + assert!(voice_matches_language("en", &v)); + assert!(voice_matches_language("english", &v)); + assert!(voice_matches_language("EN", &v)); + } + + #[test] + fn voice_matches_language_zh_matches_chinese() { + let v = VoiceInfo::builder().locale("zh").build(); + assert!(voice_matches_language("zh", &v)); + assert!(voice_matches_language("chinese", &v)); + } + + #[test] + fn voice_matches_language_mismatch() { + let v = VoiceInfo::builder().locale("en").build(); + assert!(!voice_matches_language("zh", &v)); + assert!(!voice_matches_language("ja", &v)); + } + + #[test] + fn voice_matches_language_no_locale() { + let v = VoiceInfo::builder().build(); + assert!(!voice_matches_language("en", &v)); + } + + #[test] + fn lookup_default_voice_known_and_unknown() { + assert_eq!(lookup_default_voice("Rachel"), Some("21m00Tcm4TlvDq8ikWAM")); + assert_eq!(lookup_default_voice("Adam"), Some("pNInz6obpgDQGcFmaJgB")); + assert!(lookup_default_voice("Nonexistent").is_none()); + } + + // ============ speech(mockito) ============ + + #[tokio::test] + async fn speech_returns_binary_audio() { + let mut server = mockito::Server::new_async().await; + let audio = b"\x49\x44\x33\x03\x00\x00\x00\x00\x00\x00"; // 伪 mp3 头 + server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .match_header("xi-api-key", "test-key") + .match_header("accept", "audio/mpeg") + .with_status(200) + .with_header("content-type", "audio/mpeg") + .with_body(audio) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "你好世界", "Rachel").build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.audio_data.as_deref(), Some(audio.as_ref())); + assert_eq!(result.content_type, "audio/mpeg"); + assert_eq!(result.format, "mp3"); + assert_eq!(result.model.as_deref(), Some("eleven_multilingual_v2")); + } + + #[tokio::test] + async fn speech_with_voice_id_directly() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/CustomVoiceId123") + .with_status(200) + .with_body(b"audio-bytes") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + // 直接传 voice_id(未知别名 → 原样) + let req = SpeechRequest::builder("eleven_turbo_v2", "hi", "CustomVoiceId123").build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.audio_data.as_deref(), Some(b"audio-bytes".as_ref())); + } + + #[tokio::test] + async fn speech_appends_output_format_query_for_non_mp3() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .match_query(mockito::Matcher::UrlEncoded( + "output_format".into(), + "wav".into(), + )) + .match_header("accept", "audio/wav") + .with_status(200) + .with_body(b"wav-bytes") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel") + .response_format("wav") + .build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.format, "wav"); + assert_eq!(result.content_type, "audio/wav"); + mock.assert_async().await; + } + + #[tokio::test] + async fn speech_no_output_format_query_for_mp3() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .match_query(mockito::Matcher::Missing) + .with_status(200) + .with_body(b"mp3-bytes") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let _ = adapter.speech(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn speech_sends_voice_settings_from_extra() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .match_body(mockito::Matcher::JsonString( + serde_json::json!({ + "text": "hi", + "model_id": "eleven_multilingual_v2", + "voice_settings": { + "stability": 0.5, + "similarity_boost": 0.75, + "use_speaker_boost": true + } + }) + .to_string(), + )) + .with_status(200) + .with_body(b"audio") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel") + .extra("stability", 0.5) + .extra("similarity_boost", 0.75) + .extra("use_speaker_boost", true) + .build(); + let _ = adapter.speech(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn speech_uses_default_model_when_empty() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .match_body(mockito::Matcher::JsonString( + serde_json::json!({ + "text": "hi", + "model_id": DEFAULT_MODEL + }) + .to_string(), + )) + .with_status(200) + .with_body(b"audio") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("", "hi", "Rachel").build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.model.as_deref(), Some(DEFAULT_MODEL)); + mock.assert_async().await; + } + + #[tokio::test] + async fn speech_error_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .with_status(401) + .with_body("{\"detail\":\"invalid api key\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("bad-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn speech_error_429_returns_rate_limit() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .with_status(429) + .with_body("{\"detail\":\"quota exceeded\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn speech_error_400_returns_validation() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .with_status(400) + .with_body("{\"detail\":\"text too long\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn speech_error_404_returns_voice_not_available() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/NonexistentVoice") + .with_status(404) + .with_body("{\"detail\":\"voice not found\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = + SpeechRequest::builder("eleven_multilingual_v2", "hi", "NonexistentVoice").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::VoiceNotAvailable { .. })); + } + + #[tokio::test] + async fn speech_error_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", "/text-to-speech/21m00Tcm4TlvDq8ikWAM") + .with_status(500) + .with_body("internal error") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[tokio::test] + async fn speech_without_api_key_returns_validation() { + let adapter = make_adapter(Some("https://example.com".into()), None); + let req = SpeechRequest::builder("eleven_multilingual_v2", "hi", "Rachel").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + // ============ list_voices(mockito) ============ + + #[tokio::test] + async fn list_voices_fetches_and_parses() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!({ + "voices": [ + { + "voice_id": "21m00Tcm4TlvDq8ikWAM", + "name": "Rachel", + "labels": {"gender": "female", "language": "english", "accent": "american"}, + "category": "premade" + }, + { + "voice_id": "pNInz6obpgDQGcFmaJgB", + "name": "Adam", + "labels": {"gender": "male", "language": "english"}, + "category": "premade" + } + ] + }); + server + .mock("GET", VOICES_PATH) + .match_header("xi-api-key", "test-key") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let voices = adapter.list_voices(None).await.unwrap(); + assert_eq!(voices.len(), 2); + assert_eq!(voices[0].voice_id.as_deref(), Some("21m00Tcm4TlvDq8ikWAM")); + assert_eq!(voices[0].name.as_deref(), Some("Rachel")); + assert_eq!(voices[0].gender.as_deref(), Some("Female")); + assert_eq!(voices[0].locale.as_deref(), Some("en")); + assert_eq!(voices[1].gender.as_deref(), Some("Male")); + } + + #[tokio::test] + async fn list_voices_filters_by_language() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!({ + "voices": [ + {"voice_id": "v1", "name": "EnglishVoice", "labels": {"gender": "female", "language": "english"}}, + {"voice_id": "v2", "name": "ChineseVoice", "labels": {"gender": "male", "language": "chinese"}}, + {"voice_id": "v3", "name": "JapaneseVoice", "labels": {"gender": "female", "language": "japanese"}} + ] + }); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let en = adapter.list_voices(Some("en")).await.unwrap(); + assert_eq!(en.len(), 1); + assert_eq!(en[0].voice_id.as_deref(), Some("v1")); + let zh = adapter.list_voices(Some("zh")).await.unwrap(); + assert_eq!(zh.len(), 1); + assert_eq!(zh[0].voice_id.as_deref(), Some("v2")); + let all = adapter.list_voices(None).await.unwrap(); + assert_eq!(all.len(), 3); + } + + #[tokio::test] + async fn list_voices_error_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", VOICES_PATH) + .with_status(401) + .with_body("invalid key") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("bad-key")); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_voices_error_429_returns_rate_limit() { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", VOICES_PATH) + .with_status(429) + .with_body("slow down") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn list_voices_without_api_key_returns_validation() { + let adapter = make_adapter(Some("https://example.com".into()), None); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + // ============ recommend_voices(默认实现,基于 list_voices) ============ + + #[tokio::test] + async fn recommend_voices_filters_by_gender_and_limit() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!({ + "voices": [ + {"voice_id": "v1", "name": "FemaleEn1", "labels": {"gender": "female", "language": "english"}}, + {"voice_id": "v2", "name": "MaleEn", "labels": {"gender": "male", "language": "english"}}, + {"voice_id": "v3", "name": "FemaleEn2", "labels": {"gender": "female", "language": "english"}} + ] + }); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + // 语言 + 性别过滤 + let female_en = adapter + .recommend_voices(Some("en"), Some("Female"), 10) + .await + .unwrap(); + assert_eq!(female_en.len(), 2); + for v in &female_en { + assert_eq!(v.gender.as_deref(), Some("Female")); + } + // limit 截断 + let limited = adapter.recommend_voices(Some("en"), None, 1).await.unwrap(); + assert_eq!(limited.len(), 1); + } + + #[tokio::test] + async fn recommend_voices_no_gender_returns_all_up_to_limit() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!({ + "voices": [ + {"voice_id": "v1", "name": "a", "labels": {"gender": "female", "language": "english"}}, + {"voice_id": "v2", "name": "b", "labels": {"gender": "male", "language": "english"}} + ] + }); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let all = adapter.recommend_voices(None, None, 10).await.unwrap(); + assert_eq!(all.len(), 2); + } +} From 5c2d6591e4fe251812f2ac297fd7867035bf8c2e Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:47:37 +0800 Subject: [PATCH 43/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20cartesia=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/cartesia.rs | 1416 +++++++++++++++++ 1 file changed, 1416 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/cartesia.rs diff --git a/crates/aibridge-core/src/adapters/cartesia.rs b/crates/aibridge-core/src/adapters/cartesia.rs new file mode 100644 index 0000000..44566c3 --- /dev/null +++ b/crates/aibridge-core/src/adapters/cartesia.rs @@ -0,0 +1,1416 @@ +//! Cartesia 适配器(Sonic 超低延迟 TTS,独立协议) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/audio_adapters.py` 的 `CartesiaAdapter`。 +//! +//! Cartesia 是新一代超低延迟语音合成平台,Sonic 模型以实时性和自然度著称 +//!(首包延迟 <200ms),支持多语言、情感控制与声音克隆。 +//! +//! ## 协议 +//! +//! - 认证:`X-API-Key` header + `Cartesia-Version: 2024-06-10`(非 Bearer) +//! - TTS 合成:`POST /v1/tts/bytes`,请求体 JSON,响应为二进制音频 +//! - 音色列表:`GET /tts/voices`,响应 JSON 数组(Python v1 未实现,v2 新增) +//! - 模型列表:硬编码(sonic-2 / sonic-english / sonic-multilingual) +//! +//! ## 特性 +//! +//! - `requires_api_key = true`(需 Cartesia API Key) +//! - 音色支持名称(如 "Chinese Woman")或 voice_id(UUID);名称经内置 +//! `DEFAULT_VOICES` 表解析为 ID,`extra.voice_id` 优先级最高 +//! - 输出格式:mp3(默认)/ wav / ogg / opus / raw / flac,通过 `response_format` +//! 或 `extra.output_format`(完整 dict)控制 +//! - `extra` 透传:language / speed / emotion / voice_embedding / continue / add_timestamps +//! - 音色列表带缓存,避免重复网络拉取 + +use std::collections::HashMap; + +use async_trait::async_trait; +use serde::Deserialize; +use serde_json::{json, Value}; +use tokio::sync::Mutex; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::audio::{SpeechRequest, SpeechResult}; +use crate::model::common::{ModelInfo, ModelType, VoiceInfo}; + +// ==================== 常量 ==================== + +/// Provider 类型标识 +const PROVIDER_TYPE: &str = "cartesia"; +/// Provider 显示名称 +const PROVIDER_NAME: &str = "Cartesia"; +/// 默认 API 基地址 +const DEFAULT_API_BASE: &str = "https://api.cartesia.ai"; +/// 默认 API 版本(Cartesia-Version header 值) +const DEFAULT_API_VERSION: &str = "2024-06-10"; +/// TTS 合成端点路径(对应 Python v1 `/v1/tts/bytes`) +const TTS_PATH: &str = "/v1/tts/bytes"; +/// 音色列表端点路径(Cartesia 官方 `/tts/voices`) +const VOICES_PATH: &str = "/tts/voices"; +/// 默认音色(voice 为空时兜底,对应 Python v1 默认值) +const DEFAULT_VOICE: &str = "Generic Woman"; + +/// 内置音色名称 → voice_id 映射表 +/// +/// 对应 Python v1 `CartesiaAdapter.DEFAULT_VOICES`。Cartesia 内置音色的名称 +/// 可直接作为 `speech` 的 `voice` 参数,本表将其解析为 UUID 形式的 voice_id。 +/// 传入已是 UUID 的 voice_id 时原样透传。 +fn lookup_default_voice(name: &str) -> Option<&'static str> { + let v = match name { + "Barbershop Man" => "a0e99841-438c-4a64-b679-ae501e7d6091", + "Miles" => "b27adc08-7b68-4a70-8bd3-4d7f4a4f91e8", + "Cali Bot" => "a3d1d9d3-5de0-4c8b-8325-04f0e08d0e8a", + "Customer Support Lady" => "297b6c1f-3135-403e-a9c7-28a7a4dcd690", + "Doc Brown" => "f114a467-c40a-4db8-964d-8aca59832a2f", + "Generic Man" => "c8e59920-61aa-494c-8cfe-9730e0ce4c65", + "Generic Woman" => "b2b21a4e-7f85-4905-a50e-70c04fb7cd9e", + "Helpdesk Woman" => "70ca091e-a631-4b0c-a08b-9c930ffb0a8e", + "Japanese Woman" => "8f2d130a-3a73-4d4e-bab7-d0c999371c71", + "Classy British Lady" => "56553793-917c-41a4-a057-750be38b4b61", + "Merchant" => "50d0be50-c6b1-4b10-a4dd-e10588f01b06", + "Movie Guy" => "248be419-c632-4f23-adf1-5324ed7dbf1d", + "New York Guy" => "01d31f54-303d-480a-9c27-22b9c9f3c85e", + "News Lady" => "bf9923ea-7a3f-4f62-817d-c38d59f82f33", + "Nurse" => "660569f7-267b-4240-9b74-0a4bb4dce2e5", + "Polite Man" => "156fb8d2-335b-4950-9cb3-a2d33bef8830", + "Salesman" => "87748186-23bb-4158-a1eb-332911b0b706", + "Southern Woman" => "5d1b15b6-65c6-4f3c-a349-2f2fcc5c13d1", + "The Don" => "820a3788-2b37-4d21-847a-b65d8a68c99a", + "The Laughing Guy" => "0a7e6a33-0006-447c-a723-133a0a5b09c8", + "Chemistry Professor" => "694f9389-aac1-45b6-b726-9d9369183238", + "Chinese Woman" => "e90c66b3-51bd-4906-a51a-469152d083b1", + "Sharon" => "e00d0e5a-87e7-469a-9023-d510c6f5f970", + "Competitive Podcaster" => "8c4a4d43-d33d-4a80-a9a6-a79ce1e39a69", + _ => return None, + }; + Some(v) +} + +// ==================== Cartesia 原始音色反序列化结构 ==================== + +/// Cartesia `/tts/voices` 返回的单个音色项(原始 JSON 结构) +/// +/// 字段名与 Cartesia 服务端返回的 JSON 一致。`gender` 字段 Cartesia 官方 +/// 当前不返回,保留用于兼容;缺失时由 `infer_gender` 按名称启发式推断。 +/// 其余字段(description / metadata 等)由 serde 忽略。 +#[derive(Debug, Deserialize)] +struct CartesiaVoiceRaw { + /// 音色 ID(UUID) + id: String, + /// 音色名称 + name: String, + /// 语言代码(如 "zh" / "en",部分音色为 null) + #[serde(default)] + language: Option, + /// 性别(Cartesia 官方未返回,保留兼容) + #[serde(default)] + gender: Option, +} + +// ==================== CartesiaAdapter ==================== + +/// Cartesia 适配器 +/// +/// 持有 HTTP 客户端、API key / 版本 / 基地址,以及音色列表缓存。 +/// speech 合成与 list_voices 均为 per-request HTTP 调用(无长连接)。 +pub struct CartesiaAdapter { + /// Provider 配置(保留供未来扩展,当前字段已在构造时提取) + #[allow(dead_code)] + config: ProviderConfig, + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// API 基地址(默认 `https://api.cartesia.ai`,可由 config.base_url 覆盖) + api_base: String, + /// API 版本(Cartesia-Version header 值,默认 `2024-06-10`) + api_version: String, + /// API Key(X-API-Key header 值) + api_key: String, + /// 音色列表缓存(避免重复网络拉取) + voices_cache: Mutex>>, +} + +impl CartesiaAdapter { + /// 创建 Cartesia 适配器 + /// + /// - `config.base_url` 为空时用 `DEFAULT_API_BASE` + /// - `config.api_version` 为空时用 `DEFAULT_API_VERSION` + /// - `config.api_key` 为 None 时用空串(调用时 API 会返 401) + pub fn new(config: ProviderConfig) -> Result { + let api_base = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_BASE.to_string()); + let api_version = config + .api_version + .clone() + .filter(|v| !v.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_VERSION.to_string()); + let api_key = config.api_key.clone().unwrap_or_default(); + let opts = ClientOptions::builder().timeout(config.timeout).build(); + let http = HttpClient::new(&opts)?; + Ok(Self { + config, + http, + api_base, + api_version, + api_key, + voices_cache: Mutex::new(None), + }) + } + + // ==================== 纯函数(协议构造/解析,可单测) ==================== + + /// 解析 voice_id:`extra.voice_id` > `DEFAULT_VOICES` 名称查找 > 原样透传 + /// + /// 对应 Python v1 `CartesiaAdapter._get_voice_id` + `kwargs.get("voice_id")`。 + /// voice 为空时兜底为 `DEFAULT_VOICE`("Generic Woman")。 + fn resolve_voice_id(req: &SpeechRequest) -> String { + if let Some(vid) = req.extra.get("voice_id").and_then(|v| v.as_str()) { + if !vid.is_empty() { + return vid.to_string(); + } + } + let voice = req.voice.primary().unwrap_or(""); + let voice = if voice.is_empty() { + DEFAULT_VOICE + } else { + voice + }; + lookup_default_voice(voice).unwrap_or(voice).to_string() + } + + /// 构造 `voice` 字段:有 `extra.voice_embedding` 时走克隆模式,否则走 id 模式 + /// + /// 对应 Python v1:默认 `{"mode":"id","id":voice_id}`,有 embedding 时覆盖为 + /// `{"mode":"embedding","embedding":[...]}`。 + fn build_voice(voice_id: &str, extra: &HashMap) -> Value { + if let Some(embedding) = extra.get("voice_embedding") { + return json!({"mode": "embedding", "embedding": embedding}); + } + json!({"mode": "id", "id": voice_id}) + } + + /// 构造 `output_format` 字段 + /// + /// 优先用 `extra.output_format`(完整 dict);否则由 `response_format` 构建: + /// mp3 → `{container, bit_rate:128000, sample_rate:44100}`, + /// wav → `{container, encoding:pcm_f32le, sample_rate:24000}`, + /// 其余 → `{container}`。 + fn build_output_format(response_format: &str, extra: &HashMap) -> Value { + if let Some(of) = extra.get("output_format") { + if of.is_object() { + return of.clone(); + } + } + let container = if response_format.is_empty() { + "mp3" + } else { + response_format + } + .to_lowercase(); + match container.as_str() { + "mp3" => json!({"container": "mp3", "bit_rate": 128000, "sample_rate": 44100}), + "wav" => json!({"container": "wav", "encoding": "pcm_f32le", "sample_rate": 24000}), + _ => json!({"container": container}), + } + } + + /// 从 output_format 提取 container(默认 mp3) + fn container_of(output_format: &Value) -> String { + output_format + .get("container") + .and_then(|v| v.as_str()) + .unwrap_or("mp3") + .to_string() + } + + /// container → Accept header 值(MIME 类型) + /// + /// 对应 Python v1 `accept_map`。 + fn accept_type(container: &str) -> &'static str { + match container { + "mp3" => "audio/mpeg", + "wav" => "audio/wav", + "ogg" => "audio/ogg", + "opus" => "audio/ogg;codec=opus", + "raw" | "pcm" => "audio/pcm", + "flac" => "audio/flac", + _ => "audio/mpeg", + } + } + + /// 构造 speech 请求体 + /// + /// 对应 Python v1 `CartesiaAdapter.speech` 的 payload 构造。可选字段 + ///(language / speed / emotions / continue / add_timestamps)仅在有值时加入。 + fn build_speech_payload(req: &SpeechRequest, voice_id: &str) -> Value { + let output_format = Self::build_output_format(&req.response_format, &req.extra); + let voice = Self::build_voice(voice_id, &req.extra); + let mut payload = json!({ + "model_id": req.model, + "transcript": req.input, + "voice": voice, + "output_format": output_format, + }); + + // language(extra.language) + if let Some(lang) = req.extra.get("language").and_then(|v| v.as_str()) { + payload["language"] = json!(lang); + } + + // speed(req.speed 优先,回退 extra.speed) + if let Some(speed) = req.speed { + payload["speed"] = json!(speed); + } else if let Some(speed) = req.extra.get("speed") { + payload["speed"] = speed.clone(); + } + + // emotions(req.emotion → [emotion];或 extra.emotions 列表;或 extra.emotion) + if let Some(emotion) = req.emotion.as_ref() { + payload["emotions"] = json!([emotion]); + } else if let Some(emotions) = req.extra.get("emotions") { + payload["emotions"] = emotions.clone(); + } else if let Some(emotion) = req.extra.get("emotion") { + payload["emotions"] = json!([emotion]); + } + + // continue / continuation + let cont = req + .extra + .get("continue") + .and_then(|v| v.as_bool()) + .unwrap_or(false) + || req + .extra + .get("continuation") + .and_then(|v| v.as_bool()) + .unwrap_or(false); + if cont { + payload["continue"] = json!(true); + } + + // add_timestamps + if req + .extra + .get("add_timestamps") + .and_then(|v| v.as_bool()) + .unwrap_or(false) + { + payload["add_timestamps"] = json!(true); + } + + payload + } + + /// 把 Cartesia 原始音色转为统一 `VoiceInfo` + /// + /// gender 优先用服务端返回值,缺失时按名称启发式推断。 + fn convert_voice(raw: CartesiaVoiceRaw) -> VoiceInfo { + let gender = raw.gender.or_else(|| infer_gender(&raw.name)); + VoiceInfo { + voice_id: Some(raw.id), + short_name: Some(raw.name.clone()), + name: Some(raw.name), + locale: raw.language, + gender, + extra: HashMap::new(), + } + } + + /// 拉取全量音色列表(带缓存) + /// + /// 缓存未命中时 HTTP GET `/tts/voices`,缓存后直接返回(language 过滤在 + /// `list_voices` 做)。 + async fn list_voices_raw(&self) -> Result> { + // 先检查缓存 + { + let cache = self.voices_cache.lock().await; + if let Some(ref voices) = *cache { + return Ok(voices.clone()); + } + } + let url = format!( + "{base}{path}", + base = self.api_base.trim_end_matches('/'), + path = VOICES_PATH + ); + let resp = self + .http + .inner() + .get(&url) + .header("X-API-Key", &self.api_key) + .header("Cartesia-Version", &self.api_version) + .send() + .await + .map_err(map_reqwest_error)?; + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_cartesia_error(status_code, &body)); + } + let raw_voices: Vec = resp.json().await.map_err(AibridgeError::from)?; + let voices: Vec = raw_voices.into_iter().map(Self::convert_voice).collect(); + + // 写缓存 + let mut cache = self.voices_cache.lock().await; + *cache = Some(voices.clone()); + Ok(voices) + } +} + +#[async_trait] +impl Adapter for CartesiaAdapter { + fn provider_type(&self) -> &str { + PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::AudioSpeech); + caps.insert(Capabilities::ListVoices); + caps + } + + /// Cartesia 需要 API Key(X-API-Key header) + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new() 已构造,无惰性资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // 无长连接资源需释放 + Ok(()) + } + + /// 文字转语音(Cartesia Sonic TTS) + /// + /// POST `/v1/tts/bytes`,请求体 JSON,响应二进制音频。voice 列表取首个 + ///(Cartesia 不支持音色降级,与 Edge TTS 不同)。 + async fn speech(&self, req: SpeechRequest) -> Result { + let voice_id = Self::resolve_voice_id(&req); + let payload = Self::build_speech_payload(&req, &voice_id); + // 从 payload 取 output_format 派生 container/Accept,保证与请求体一致 + let output_format = payload + .get("output_format") + .cloned() + .unwrap_or_else(|| json!({"container": "mp3"})); + let container = Self::container_of(&output_format); + let accept = Self::accept_type(&container); + let url = format!( + "{base}{path}", + base = self.api_base.trim_end_matches('/'), + path = TTS_PATH + ); + + let resp = self + .http + .inner() + .post(&url) + .header("X-API-Key", &self.api_key) + .header("Cartesia-Version", &self.api_version) + .header("Accept", accept) + .json(&payload) + .send() + .await + .map_err(map_reqwest_error)?; + + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_cartesia_error(status_code, &body)); + } + + let audio_data = resp.bytes().await.map_err(map_reqwest_error)?.to_vec(); + + // 空音频检测:TTS 不应返回空音频,判为服务端临时问题(可重试) + if audio_data.is_empty() { + return Err(AibridgeError::service_unavailable( + "Cartesia 返回空音频(可能是限流或服务端抖动),可重试", + )); + } + + Ok(SpeechResult { + audio_data: Some(audio_data), + audio_url: None, + audio_base64: None, + content_type: accept.to_string(), + format: container, + duration: None, + model: Some(if req.model.is_empty() { + PROVIDER_TYPE.to_string() + } else { + req.model.clone() + }), + }) + } + + /// 列出可用音色(带缓存,按 language 前缀过滤) + async fn list_voices(&self, language: Option<&str>) -> Result> { + let voices = self.list_voices_raw().await?; + match language { + Some(lang) if !lang.is_empty() => Ok(voices + .into_iter() + .filter(|v| { + v.locale + .as_deref() + .map(|l| l.starts_with(lang)) + .unwrap_or(false) + }) + .collect()), + _ => Ok(voices), + } + } + + /// 列出 Cartesia 模型(硬编码,不实时拉取) + async fn list_models(&self, filter: Option) -> Result> { + let models = vec![ + ModelInfo { + id: "sonic-2".into(), + name: "Cartesia Sonic 2".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Cartesia Sonic 2 超低延迟多语言 TTS 模型(首包 <200ms)".into()), + created: None, + }, + ModelInfo { + id: "sonic-english".into(), + name: "Cartesia Sonic English".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Cartesia Sonic 英文专用 TTS 模型".into()), + created: None, + }, + ModelInfo { + id: "sonic-multilingual".into(), + name: "Cartesia Sonic Multilingual".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_speech".into()], + max_tokens: None, + supports_streaming: false, + description: Some("Cartesia Sonic 多语言 TTS 模型".into()), + created: None, + }, + ]; + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 错误映射 ==================== + +/// 将 Cartesia HTTP 错误响应映射为 AibridgeError +/// +/// 对应 Python v1 `CartesiaAdapter._handle_error`。映射规则(与任务规格一致): +/// - 401/403 → Authentication(X-API-Key 无效) +/// - 429 → RateLimit(限流或配额耗尽) +/// - 400/422 → Validation(参数校验错误,携带 details) +/// - 404 → VoiceNotAvailable(voice_id 或模型不存在) +/// - 5xx → Api(服务端错误) +/// - 其余 4xx → Api +fn map_cartesia_error(status: u16, body: &str) -> AibridgeError { + let message = parse_cartesia_error_message(body, status); + match status { + 401 | 403 => AibridgeError::authentication(format!( + "Cartesia 认证失败(X-API-Key 无效或缺失): {message}" + )), + 429 => AibridgeError::rate_limit(format!("Cartesia 限流或配额耗尽: {message}")), + 400 | 422 => AibridgeError::validation_with_details( + format!("Cartesia 参数校验错误: {message}"), + serde_json::json!({"status": status, "response": body}), + ), + 404 => { + AibridgeError::voice_not_available(format!("Cartesia voice_id 或模型不存在: {message}")) + } + s if s >= 500 => AibridgeError::api(s, format!("Cartesia 服务端错误 ({s}): {message}")), + s => AibridgeError::api(s, format!("Cartesia HTTP {s}: {message}")), + } +} + +/// 从 Cartesia 错误响应体解析错误消息 +/// +/// 尝试解析 JSON 的 `message` 或 `error` 字段(error 可为字符串或 +/// `{"message": "..."}` 对象);解析失败时截取前 200 字符或回退 `HTTP {status}`。 +fn parse_cartesia_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + if let Some(m) = v.get("message").and_then(|m| m.as_str()) { + return m.to_string(); + } + if let Some(error) = v.get("error") { + if let Some(s) = error.as_str() { + return s.to_string(); + } + if let Some(m) = error.get("message").and_then(|m| m.as_str()) { + return m.to_string(); + } + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + body.chars().take(200).collect() + } +} + +// ==================== 辅助函数 ==================== + +/// 将 reqwest::Error 映射为 AibridgeError(超时 → Timeout,其余 → Network) +fn map_reqwest_error(e: reqwest::Error) -> AibridgeError { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } +} + +/// 按音色名称启发式推断性别 +/// +/// Cartesia `/tts/voices` 不返回 gender 字段,本函数按名称关键词推断: +/// 含 "woman"/"lady"/"girl" → Female;含 "man"/"guy"/"boy" → Male; +/// 其余 → None(未知)。注意先判 "woman"(含 "man" 子串)再判 "man"。 +fn infer_gender(name: &str) -> Option { + let lower = name.to_lowercase(); + if lower.contains("woman") || lower.contains("lady") || lower.contains("girl") { + Some("Female".to_string()) + } else if lower.contains("man") || lower.contains("guy") || lower.contains("boy") { + Some("Male".to_string()) + } else { + None + } +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::audio::TranscribeRequest; + use crate::model::chat::ChatRequest; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + use std::collections::HashMap; + + /// 构造测试用适配器(base_url 指向 mockito server,None 用默认基地址) + fn make_adapter(base_url: Option) -> CartesiaAdapter { + let mut opts = ClientOptions::builder().api_key("test-key"); + if let Some(u) = base_url { + opts = opts.base_url(u); + } + let config = ProviderConfig::from_options(PROVIDER_TYPE, opts.build()); + CartesiaAdapter::new(config).expect("构造 CartesiaAdapter 失败") + } + + // ============ 基本属性 ============ + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter(None); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_speech_and_list_voices() { + let adapter = make_adapter(None); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::AudioSpeech)); + assert!(caps.contains(&Capabilities::ListVoices)); + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + } + + #[test] + fn provider_type_and_name() { + let adapter = make_adapter(None); + assert_eq!(adapter.provider_type(), "cartesia"); + assert_eq!(adapter.provider_name(), "Cartesia"); + } + + #[tokio::test] + async fn list_models_returns_sonic_models() { + let adapter = make_adapter(None); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 3); + assert!(models.iter().any(|m| m.id == "sonic-2")); + assert!(models.iter().any(|m| m.id == "sonic-english")); + assert!(models.iter().any(|m| m.id == "sonic-multilingual")); + for m in &models { + assert_eq!(m.model_type, ModelType::Audio); + assert_eq!(m.provider, "cartesia"); + } + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let adapter = make_adapter(None); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 3); + let chat = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat.is_empty()); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter(None); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + #[test] + fn api_version_defaults_when_unset() { + let adapter = make_adapter(None); + assert_eq!(adapter.api_version, DEFAULT_API_VERSION); + } + + // ============ 不支持能力(默认实现) ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter(None); + let req = ChatRequest::builder("m", vec![]).build(); + assert!(matches!( + adapter.chat(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn image_generate_returns_unsupported() { + let adapter = make_adapter(None); + let req = ImageRequest::builder("m", "p").build(); + assert!(matches!( + adapter.image_generate(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter(None); + let req = VideoRequest::builder("m", "p").build(); + assert!(matches!( + adapter.video_create(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn embed_returns_unsupported() { + let adapter = make_adapter(None); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + assert!(matches!( + adapter.embed(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn transcribe_returns_unsupported() { + let adapter = make_adapter(None); + let req = TranscribeRequest::builder("m", FileInput::path("/tmp/a.mp3")).build(); + assert!(matches!( + adapter.transcribe(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + // ============ resolve_voice_id ============ + + #[test] + fn resolve_voice_id_name_lookup() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + assert_eq!( + CartesiaAdapter::resolve_voice_id(&req), + "e90c66b3-51bd-4906-a51a-469152d083b1" + ); + let req = SpeechRequest::builder("sonic-2", "hi", "Barbershop Man").build(); + assert_eq!( + CartesiaAdapter::resolve_voice_id(&req), + "a0e99841-438c-4a64-b679-ae501e7d6091" + ); + } + + #[test] + fn resolve_voice_id_passthrough_uuid() { + // 非 known 名称(UUID 形式)原样透传 + let req = SpeechRequest::builder("sonic-2", "hi", "custom-voice-uuid").build(); + assert_eq!(CartesiaAdapter::resolve_voice_id(&req), "custom-voice-uuid"); + } + + #[test] + fn resolve_voice_id_empty_falls_back_to_default() { + let req = SpeechRequest::builder("sonic-2", "hi", "").build(); + assert_eq!( + CartesiaAdapter::resolve_voice_id(&req), + "b2b21a4e-7f85-4905-a50e-70c04fb7cd9e" // Generic Woman + ); + } + + #[test] + fn resolve_voice_id_extra_voice_id_takes_priority() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("voice_id", "override-uuid") + .build(); + assert_eq!(CartesiaAdapter::resolve_voice_id(&req), "override-uuid"); + } + + #[test] + fn resolve_voice_id_empty_extra_voice_id_ignored() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("voice_id", "") + .build(); + assert_eq!( + CartesiaAdapter::resolve_voice_id(&req), + "e90c66b3-51bd-4906-a51a-469152d083b1" + ); + } + + // ============ build_output_format / accept_type / container_of ============ + + #[test] + fn build_output_format_default_mp3() { + let of = CartesiaAdapter::build_output_format("mp3", &HashMap::new()); + assert_eq!(of["container"], "mp3"); + assert_eq!(of["bit_rate"], 128000); + assert_eq!(of["sample_rate"], 44100); + } + + #[test] + fn build_output_format_wav() { + let of = CartesiaAdapter::build_output_format("wav", &HashMap::new()); + assert_eq!(of["container"], "wav"); + assert_eq!(of["encoding"], "pcm_f32le"); + assert_eq!(of["sample_rate"], 24000); + } + + #[test] + fn build_output_format_empty_defaults_mp3() { + let of = CartesiaAdapter::build_output_format("", &HashMap::new()); + assert_eq!(of["container"], "mp3"); + } + + #[test] + fn build_output_format_extra_overrides() { + let mut extra = HashMap::new(); + extra.insert( + "output_format".to_string(), + json!({"container": "raw", "sample_rate": 16000}), + ); + let of = CartesiaAdapter::build_output_format("mp3", &extra); + assert_eq!(of["container"], "raw"); + assert_eq!(of["sample_rate"], 16000); + } + + #[test] + fn accept_type_mapping() { + assert_eq!(CartesiaAdapter::accept_type("mp3"), "audio/mpeg"); + assert_eq!(CartesiaAdapter::accept_type("wav"), "audio/wav"); + assert_eq!(CartesiaAdapter::accept_type("ogg"), "audio/ogg"); + assert_eq!(CartesiaAdapter::accept_type("opus"), "audio/ogg;codec=opus"); + assert_eq!(CartesiaAdapter::accept_type("raw"), "audio/pcm"); + assert_eq!(CartesiaAdapter::accept_type("pcm"), "audio/pcm"); + assert_eq!(CartesiaAdapter::accept_type("flac"), "audio/flac"); + assert_eq!(CartesiaAdapter::accept_type("unknown"), "audio/mpeg"); + } + + #[test] + fn container_of_extracts_or_defaults() { + assert_eq!( + CartesiaAdapter::container_of(&json!({"container": "wav"})), + "wav" + ); + assert_eq!(CartesiaAdapter::container_of(&json!({})), "mp3"); + } + + // ============ build_voice / build_speech_payload ============ + + #[test] + fn build_voice_id_mode() { + let v = CartesiaAdapter::build_voice("vid-123", &HashMap::new()); + assert_eq!(v["mode"], "id"); + assert_eq!(v["id"], "vid-123"); + } + + #[test] + fn build_voice_embedding_mode() { + let mut extra = HashMap::new(); + extra.insert("voice_embedding".to_string(), json!([0.1, 0.2, 0.3])); + let v = CartesiaAdapter::build_voice("vid-123", &extra); + assert_eq!(v["mode"], "embedding"); + assert_eq!(v["embedding"][2], 0.3); + assert!(v.get("id").is_none()); + } + + #[test] + fn build_speech_payload_minimal() { + let req = SpeechRequest::builder("sonic-2", "你好", "Chinese Woman").build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["model_id"], "sonic-2"); + assert_eq!(payload["transcript"], "你好"); + assert_eq!(payload["voice"]["mode"], "id"); + assert_eq!(payload["voice"]["id"], "voice-uuid"); + assert_eq!(payload["output_format"]["container"], "mp3"); + // 无可选字段时不应出现 + assert!(payload.get("language").is_none()); + assert!(payload.get("speed").is_none()); + assert!(payload.get("emotions").is_none()); + assert!(payload.get("continue").is_none()); + assert!(payload.get("add_timestamps").is_none()); + } + + #[test] + fn build_speech_payload_with_all_optional() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .speed(1.5) + .emotion("positivity") + .extra("language", "zh") + .extra("continue", true) + .extra("add_timestamps", true) + .build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["language"], "zh"); + assert_eq!(payload["speed"], 1.5); + assert_eq!(payload["emotions"], json!(["positivity"])); + assert_eq!(payload["continue"], true); + assert_eq!(payload["add_timestamps"], true); + } + + #[test] + fn build_speech_payload_embedding_overrides_voice() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("voice_embedding", json!([0.1, 0.2])) + .build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["voice"]["mode"], "embedding"); + assert!(payload["voice"].get("id").is_none()); + } + + #[test] + fn build_speech_payload_emotions_from_extra_list() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("emotions", json!(["curiosity", "surprise"])) + .build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["emotions"], json!(["curiosity", "surprise"])); + } + + #[test] + fn build_speech_payload_speed_from_extra_when_req_speed_none() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("speed", 2.0) + .build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["speed"], 2.0); + } + + #[test] + fn build_speech_payload_continuation_alias() { + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .extra("continuation", true) + .build(); + let payload = CartesiaAdapter::build_speech_payload(&req, "voice-uuid"); + assert_eq!(payload["continue"], true); + } + + // ============ infer_gender ============ + + #[test] + fn infer_gender_female_keywords() { + assert_eq!(infer_gender("Chinese Woman").as_deref(), Some("Female")); + assert_eq!(infer_gender("News Lady").as_deref(), Some("Female")); + assert_eq!(infer_gender("Business Girl").as_deref(), Some("Female")); + } + + #[test] + fn infer_gender_male_keywords() { + assert_eq!(infer_gender("Barbershop Man").as_deref(), Some("Male")); + assert_eq!(infer_gender("Movie Guy").as_deref(), Some("Male")); + assert_eq!(infer_gender("News Boy").as_deref(), Some("Male")); + } + + #[test] + fn infer_gender_woman_checked_before_man() { + // "Woman" 含 "man" 子串,必须先判 woman 返 Female + assert_eq!(infer_gender("Woman").as_deref(), Some("Female")); + assert_eq!(infer_gender("Salesman").as_deref(), Some("Male")); + } + + #[test] + fn infer_gender_none_for_neutral() { + assert!(infer_gender("Sharon").is_none()); + assert!(infer_gender("Competitive Podcaster").is_none()); + } + + // ============ convert_voice ============ + + #[test] + fn convert_voice_maps_all_fields() { + let raw = CartesiaVoiceRaw { + id: "uuid-1".into(), + name: "Chinese Woman".into(), + language: Some("zh".into()), + gender: None, + }; + let v = CartesiaAdapter::convert_voice(raw); + assert_eq!(v.voice_id.as_deref(), Some("uuid-1")); + assert_eq!(v.short_name.as_deref(), Some("Chinese Woman")); + assert_eq!(v.name.as_deref(), Some("Chinese Woman")); + assert_eq!(v.locale.as_deref(), Some("zh")); + assert_eq!(v.gender.as_deref(), Some("Female")); // 名称推断 + } + + #[test] + fn convert_voice_uses_explicit_gender_when_present() { + let raw = CartesiaVoiceRaw { + id: "uuid-2".into(), + name: "Some Voice".into(), + language: None, + gender: Some("Male".into()), + }; + let v = CartesiaAdapter::convert_voice(raw); + assert_eq!(v.gender.as_deref(), Some("Male")); // 服务端值优先,不推断 + } + + // ============ 错误解析 ============ + + #[test] + fn parse_error_message_from_message_field() { + let body = r#"{"message":"voice not found"}"#; + assert_eq!(parse_cartesia_error_message(body, 404), "voice not found"); + } + + #[test] + fn parse_error_message_from_error_string() { + let body = r#"{"error":"bad request"}"#; + assert_eq!(parse_cartesia_error_message(body, 400), "bad request"); + } + + #[test] + fn parse_error_message_from_error_object() { + let body = r#"{"error":{"message":"invalid model_id"}}"#; + assert_eq!(parse_cartesia_error_message(body, 422), "invalid model_id"); + } + + #[test] + fn parse_error_message_fallback_to_body_text() { + assert_eq!( + parse_cartesia_error_message("plain text error", 500), + "plain text error" + ); + } + + #[test] + fn parse_error_message_empty_body_falls_back_to_http_status() { + assert_eq!(parse_cartesia_error_message("", 500), "HTTP 500"); + assert_eq!(parse_cartesia_error_message(" ", 500), "HTTP 500"); + } + + #[test] + fn map_cartesia_error_401_is_authentication() { + let err = map_cartesia_error(401, r#"{"message":"invalid api key"}"#); + assert!(matches!(err, AibridgeError::Authentication { .. })); + assert!(!err.is_retryable()); + } + + #[test] + fn map_cartesia_error_403_is_authentication() { + let err = map_cartesia_error(403, "forbidden"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_cartesia_error_429_is_rate_limit() { + let err = map_cartesia_error(429, r#"{"message":"slow down"}"#); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + assert!(err.is_retryable()); + } + + #[test] + fn map_cartesia_error_400_is_validation() { + let err = map_cartesia_error(400, r#"{"message":"bad param"}"#); + assert!(matches!(err, AibridgeError::Validation { .. })); + assert!(!err.is_retryable()); + } + + #[test] + fn map_cartesia_error_422_is_validation() { + let err = map_cartesia_error(422, r#"{"error":"unprocessable"}"#); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_cartesia_error_404_is_voice_not_available() { + let err = map_cartesia_error(404, r#"{"message":"voice not found"}"#); + assert!(matches!(err, AibridgeError::VoiceNotAvailable { .. })); + } + + #[test] + fn map_cartesia_error_500_is_api_and_retryable() { + let err = map_cartesia_error(500, r#"{"message":"server error"}"#); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + assert!(err.is_retryable()); + } + + #[test] + fn map_cartesia_error_other_4xx_is_api() { + let err = map_cartesia_error(418, "teapot"); + assert!(matches!(err, AibridgeError::Api { status: 418, .. })); + assert!(!err.is_retryable()); + } + + // ============ speech(HTTP,mockito) ============ + + #[tokio::test] + async fn speech_returns_binary_audio() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .match_header("X-API-Key", "test-key") + .match_header("Cartesia-Version", "2024-06-10") + .match_header("Accept", "audio/mpeg") + .with_status(200) + .with_header("content-type", "audio/mpeg") + .with_body(b"\xFF\xFB\x90\x00\x01\x02\x03") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "你好世界", "Chinese Woman").build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!( + result.audio_data, + Some(vec![0xFF, 0xFB, 0x90, 0x00, 0x01, 0x02, 0x03]) + ); + assert_eq!(result.content_type, "audio/mpeg"); + assert_eq!(result.format, "mp3"); + assert_eq!(result.model.as_deref(), Some("sonic-2")); + } + + #[tokio::test] + async fn speech_sends_correct_payload() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", TTS_PATH) + .match_header("X-API-Key", "test-key") + .match_body(mockito::Matcher::Json(json!({ + "model_id": "sonic-2", + "transcript": "hello", + "voice": {"mode": "id", "id": "e90c66b3-51bd-4906-a51a-469152d083b1"}, + "output_format": {"container": "mp3", "bit_rate": 128000, "sample_rate": 44100} + }))) + .with_status(200) + .with_header("content-type", "audio/mpeg") + .with_body(b"\x00\x01\x02") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hello", "Chinese Woman").build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.audio_data, Some(vec![0x00, 0x01, 0x02])); + mock.assert_async().await; + } + + #[tokio::test] + async fn speech_wav_format_sets_accept_header() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .match_header("Accept", "audio/wav") + .with_status(200) + .with_header("content-type", "audio/wav") + .with_body(b"RIFFxxxxWAVE") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman") + .response_format("wav") + .build(); + let result = adapter.speech(req).await.unwrap(); + assert_eq!(result.format, "wav"); + assert_eq!(result.content_type, "audio/wav"); + } + + #[tokio::test] + async fn speech_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(401) + .with_body(r#"{"message":"Invalid API key"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn speech_429_returns_rate_limit() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(429) + .with_body(r#"{"message":"rate limit"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn speech_400_returns_validation() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(400) + .with_body(r#"{"error":"transcript too long"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn speech_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(500) + .with_body(r#"{"message":"internal error"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[tokio::test] + async fn speech_404_returns_voice_not_available() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(404) + .with_body(r#"{"message":"voice not found"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::VoiceNotAvailable { .. })); + } + + #[tokio::test] + async fn speech_empty_audio_returns_service_unavailable() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", TTS_PATH) + .with_status(200) + .with_body(b"") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = SpeechRequest::builder("sonic-2", "hi", "Chinese Woman").build(); + let err = adapter.speech(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::ServiceUnavailable { .. })); + } + + // ============ list_voices(HTTP,mockito) ============ + + #[tokio::test] + async fn list_voices_fetches_and_parses() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"id": "e90c66b3-51bd-4906-a51a-469152d083b1", "name": "Chinese Woman", "language": "zh"}, + {"id": "a0e99841-438c-4a64-b679-ae501e7d6091", "name": "Barbershop Man", "language": "en"} + ]); + server + .mock("GET", VOICES_PATH) + .match_header("X-API-Key", "test-key") + .match_header("Cartesia-Version", "2024-06-10") + .with_status(200) + .with_header("content-type", "application/json") + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let voices = adapter.list_voices(None).await.unwrap(); + assert_eq!(voices.len(), 2); + assert_eq!( + voices[0].voice_id.as_deref(), + Some("e90c66b3-51bd-4906-a51a-469152d083b1") + ); + assert_eq!(voices[0].name.as_deref(), Some("Chinese Woman")); + assert_eq!(voices[0].locale.as_deref(), Some("zh")); + assert_eq!(voices[0].gender.as_deref(), Some("Female")); // 名称推断 + assert_eq!(voices[1].gender.as_deref(), Some("Male")); // 名称推断 + } + + #[tokio::test] + async fn list_voices_filters_by_language() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"id": "v1", "name": "Chinese Woman", "language": "zh"}, + {"id": "v2", "name": "Japanese Woman", "language": "ja"}, + {"id": "v3", "name": "Barbershop Man", "language": "en"} + ]); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let zh = adapter.list_voices(Some("zh")).await.unwrap(); + assert_eq!(zh.len(), 1); + assert_eq!(zh[0].voice_id.as_deref(), Some("v1")); + let none = adapter.list_voices(Some("fr")).await.unwrap(); + assert!(none.is_empty()); + } + + #[tokio::test] + async fn list_voices_caches_second_call() { + let mut server = mockito::Server::new_async().await; + let m = server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body( + serde_json::json!([{"id": "v1", "name": "Chinese Woman", "language": "zh"}]) + .to_string(), + ) + .expect(1) // 第二次走缓存,不命中 HTTP + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let _ = adapter.list_voices(None).await.unwrap(); + let _ = adapter.list_voices(None).await.unwrap(); // 走缓存 + m.assert_async().await; + } + + #[tokio::test] + async fn list_voices_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", VOICES_PATH) + .with_status(401) + .with_body(r#"{"message":"invalid key"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn list_voices_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("GET", VOICES_PATH) + .with_status(503) + .with_body(r#"{"message":"unavailable"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let err = adapter.list_voices(None).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { .. })); + } + + // ============ recommend_voices ============ + + #[tokio::test] + async fn recommend_voices_filters_by_gender_and_limit() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"id": "v1", "name": "Chinese Woman", "language": "zh"}, + {"id": "v2", "name": "Barbershop Man", "language": "zh"}, + {"id": "v3", "name": "Japanese Woman", "language": "ja"} + ]); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + // 中文 + Female → 仅 v1 + let female_zh = adapter + .recommend_voices(Some("zh"), Some("Female"), 10) + .await + .unwrap(); + assert_eq!(female_zh.len(), 1); + assert_eq!(female_zh[0].voice_id.as_deref(), Some("v1")); + + // 中文不限性别 → v1 + v2,limit=1 → 仅 v1 + let limited = adapter.recommend_voices(Some("zh"), None, 1).await.unwrap(); + assert_eq!(limited.len(), 1); + } + + #[tokio::test] + async fn recommend_voices_male_filter() { + let mut server = mockito::Server::new_async().await; + let body = serde_json::json!([ + {"id": "v1", "name": "Chinese Woman", "language": "zh"}, + {"id": "v2", "name": "Barbershop Man", "language": "zh"}, + {"id": "v3", "name": "Movie Guy", "language": "en"} + ]); + server + .mock("GET", VOICES_PATH) + .with_status(200) + .with_body(body.to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let males = adapter + .recommend_voices(None, Some("male"), 10) + .await + .unwrap(); + assert_eq!(males.len(), 2); // v2 + v3 + } +} From ac0a3c869b3e103cf9a6a4cd95046a53e13eefa5 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 01:51:41 +0800 Subject: [PATCH 44/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20=E7=AC=AC=E4=BA=8C=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20elevenlabs=20+=20cartesia=20=E5=88=B0=E5=B7=A5?= =?UTF-8?q?=E5=8E=82?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 55 ++++++++++++++++++--- crates/aibridge-core/src/adapters/mod.rs | 6 +++ 2 files changed, 53 insertions(+), 8 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index d4734c9..94fb03f 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -17,11 +17,13 @@ use crate::adapters::aggregation_platforms::{ use crate::adapters::agnes::AgnesAdapter; use crate::adapters::anthropic::AnthropicAdapter; use crate::adapters::azure::AzureAdapter; +use crate::adapters::cartesia::CartesiaAdapter; use crate::adapters::chinese::{ DoubaoAdapter, ErnieAdapter, KimiAdapter, MiniMaxAdapter, QwenAdapter, ZhipuAdapter, }; use crate::adapters::echo::EchoAdapter; use crate::adapters::edge_tts::EdgeTtsAdapter; +use crate::adapters::elevenlabs::ElevenLabsAdapter; use crate::adapters::emerging_models::{IdeogramAdapter, LlamaAdapter, LumaAdapter}; use crate::adapters::gemini::GeminiAdapter; use crate::adapters::kling::KlingAdapter; @@ -83,9 +85,9 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "kling", // 阶段 2c 音频(已实现): "edge-tts", - // 阶段 2c 待实现: "elevenlabs", "cartesia", + // 阶段 2c 待实现: "deepgram", "assemblyai", ]; @@ -149,12 +151,14 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { // 阶段 2c 音频:别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用 // (edge-tts / edge_tts / edge 三个别名均指向 EdgeTTSAdapter,免认证) "edge-tts" | "edge_tts" | "edge" => Ok(Box::new(EdgeTtsAdapter::new(config)?)), + // (elevenlabs / eleven / 11labs 均指向 ElevenLabsAdapter) + "elevenlabs" | "eleven" | "11labs" => Ok(Box::new(ElevenLabsAdapter::new(config)?)), + // (cartesia / sonic 均指向 CartesiaAdapter) + "cartesia" | "sonic" => Ok(Box::new(CartesiaAdapter::new(config)?)), // 阶段 2c 适配器占位 - "elevenlabs" | "cartesia" | "deepgram" | "assemblyai" => { - Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2c 待实现)"), - }) - } + "deepgram" | "assemblyai" => Err(AibridgeError::ProviderNotFound { + provider: format!("{provider}(阶段 2c 待实现)"), + }), // 未知 provider _ => Err(AibridgeError::provider_not_found(format!( "{provider}(未知 provider,支持:{})", @@ -478,6 +482,39 @@ mod tests { assert_eq!(edge.provider_type(), "edge-tts"); } + #[test] + fn create_elevenlabs_returns_adapter() { + // 阶段 2c:ElevenLabsAdapter 自带 DEFAULT_API_BASE 兜底,仅需 api_key + let adapter = + create_adapter(config_for("elevenlabs")).expect("工厂应能创建 elevenlabs 适配器"); + assert_eq!(adapter.provider_type(), "elevenlabs"); + } + + #[test] + fn create_cartesia_returns_adapter() { + // 阶段 2c:CartesiaAdapter 自带 DEFAULT_API_BASE 兜底,仅需 api_key + let adapter = create_adapter(config_for("cartesia")).expect("工厂应能创建 cartesia 适配器"); + assert_eq!(adapter.provider_type(), "cartesia"); + } + + #[test] + fn create_elevenlabs_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用: + // eleven / 11labs -> elevenlabs(均指向 ElevenLabsAdapter) + let eleven = create_adapter(config_for("eleven")).expect("别名 eleven 应映射到 elevenlabs"); + assert_eq!(eleven.provider_type(), "elevenlabs"); + let labs = create_adapter(config_for("11labs")).expect("别名 11labs 应映射到 elevenlabs"); + assert_eq!(labs.provider_type(), "elevenlabs"); + } + + #[test] + fn create_cartesia_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用: + // sonic -> cartesia(指向 CartesiaAdapter) + let sonic = create_adapter(config_for("sonic")).expect("别名 sonic 应映射到 cartesia"); + assert_eq!(sonic.provider_type(), "cartesia"); + } + #[test] fn create_additional_models_aliases_map_to_main_provider_type() { // 别名对齐 Python agn/adapters/additional_models.py 末尾 register 调用: @@ -505,8 +542,8 @@ mod tests { #[test] fn create_phase2_adapter_returns_phase2_message() { - // elevenlabs 仍为阶段 2c 占位(未实现),返 ProviderNotFound - let result = create_adapter(config_for("elevenlabs")); + // deepgram 仍为阶段 2c 占位(未实现),返 ProviderNotFound + let result = create_adapter(config_for("deepgram")); if let Err(AibridgeError::ProviderNotFound { provider }) = result { assert!(provider.contains("阶段 2")); } else { @@ -519,6 +556,8 @@ mod tests { assert!(is_known_provider("echo")); assert!(is_known_provider("openai")); assert!(is_known_provider("edge-tts")); + assert!(is_known_provider("elevenlabs")); + assert!(is_known_provider("cartesia")); assert!(is_known_provider("assemblyai")); assert!(is_known_provider("kling")); // 阶段 2a 已实现 provider 应被识别 diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 5a3d71f..7851c84 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -63,3 +63,9 @@ pub mod kling; /// Edge TTS 适配器:阶段 2c 音频,免费文字转语音(免认证,WebSocket 协议) pub mod edge_tts; + +/// ElevenLabs 适配器:阶段 2c 音频,ElevenLabs TTS 文字转语音(高质量音色/多语种/克隆) +pub mod elevenlabs; + +/// Cartesia 适配器:阶段 2c 音频,Cartesia Sonic TTS 文字转语音(低延迟流式) +pub mod cartesia; From e533bff2aaaee956a8b13182dfd2d4e454f17e6d Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 02:05:07 +0800 Subject: [PATCH 45/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20deepgram=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapters/deepgram.rs | 1447 +++++++++++++++++ 1 file changed, 1447 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/deepgram.rs diff --git a/crates/aibridge-core/src/adapters/deepgram.rs b/crates/aibridge-core/src/adapters/deepgram.rs new file mode 100644 index 0000000..9afa8ce --- /dev/null +++ b/crates/aibridge-core/src/adapters/deepgram.rs @@ -0,0 +1,1447 @@ +//! Deepgram 适配器(超低延迟语音识别 ASR) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/audio_adapters.py` 的 `DeepgramAdapter`。 +//! +//! Deepgram 是全球最快的语音识别(ASR)服务之一,Nova 系列模型以超低延迟和高准确率著称。 +//! +//! ## 协议(独立协议,非 OpenAI 兼容) +//! +//! - Base URL: `https://api.deepgram.com/v1` +//! - 认证: `Authorization: Token `(非 Bearer) +//! - transcribe: `POST /listen` +//! - 音频以二进制 raw body 上传(`Content-Type` 为音频 MIME),同时通过 query param 传 +//! `model` / `smart_format` / `punctuate` / `language` / `diarize` / `profanity_filter` / +//! `utterances` / `keywords` / `detect_language` 等参数 +//! - 远程 URL 输入:不传 body,改用 `url` query param(Deepgram 原生远程音频端点) +//! - 响应: JSON,结构为 `{ results: { channels: [{ alternatives: [{ transcript, confidence, +//! words, paragraphs }] }] }, metadata: { duration, model_info: { name, language } } }` +//! - list_models: 无标准 /models 端点,保留硬编码 ASR 模型列表 +//! - 文档: https://developers.deepgram.com/reference/listen-file +//! +//! ## 错误映射 +//! +//! - 401 → Authentication(API Key/Token 无效) +//! - 403 → Authentication(权限不足或账户停用) +//! - 429 → RateLimit(限流或配额耗尽) +//! - 400/422 → Validation(请求参数错误) +//! - 5xx → Api(服务端错误) +//! - 其余 4xx → Api +//! +//! ## 特性 +//! +//! - `requires_api_key = true`(需 API Key) +//! - `capabilities` 仅 `AudioTranscribe`(transcribe + translate 均走 `transcribe()` 方法) +//! - Deepgram 不支持翻译为英文(`translate=true` 时返回 Validation 错误,与 Python v1 一致: +//! Python 版 `DeepgramAdapter` 也未实现 translate) + +use async_trait::async_trait; +use serde_json::Value; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::audio::{ + TranscribeRequest, TranscriptionResult, TranscriptionSegment, TranscriptionWord, +}; +use crate::model::common::{ModelInfo, ModelType}; +use crate::model::image::FileInput; + +// ==================== 常量 ==================== + +/// Provider 类型标识 +const PROVIDER_TYPE: &str = "deepgram"; +/// Provider 显示名称 +const PROVIDER_NAME: &str = "Deepgram"; +/// 默认 API 基地址 +const DEFAULT_API_BASE: &str = "https://api.deepgram.com/v1"; +/// 预录音频转写端点路径 +const LISTEN_PATH: &str = "/listen"; +/// 默认模型(Nova-2 通用,Deepgram 推荐默认) +const DEFAULT_MODEL: &str = "nova-2"; + +// ==================== 音频输入解析 ==================== + +/// 解析后的音频输入形式 +/// +/// Deepgram 支持两种音频输入方式: +/// - 直接上传二进制 body(Bytes/Base64/Path) +/// - 远程 URL(通过 `url` query param,不传 body) +#[derive(Debug)] +enum AudioInput { + /// 二进制音频 + Content-Type(直接作为请求 body 上传) + Bytes(Vec, String), + /// 远程 URL(用 `url` query param,不传 body) + RemoteUrl(String), +} + +// ==================== DeepgramAdapter ==================== + +/// Deepgram 适配器 +/// +/// 持有 HTTP 客户端与 API Key。所有请求按需发起(无长连接)。 +pub struct DeepgramAdapter { + /// Provider 配置 + #[allow(dead_code)] + config: ProviderConfig, + /// HTTP 客户端 + http: HttpClient, + /// API Key(`Authorization: Token ` 用) + api_key: Option, + /// API 基地址(默认 `https://api.deepgram.com/v1`,可由 config.base_url 覆盖) + api_base: String, +} + +impl DeepgramAdapter { + /// 创建 Deepgram 适配器 + /// + /// `config.base_url` 可覆盖 API 基地址(主要用于测试指向 mock server), + /// 为空时用 `DEFAULT_API_BASE`。API Key 从 `config.api_key` 取。 + pub fn new(config: ProviderConfig) -> Result { + let api_base = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_BASE.to_string()); + let api_key = config.api_key.clone(); + let opts = ClientOptions::builder().timeout(config.timeout).build(); + let http = HttpClient::new(&opts)?; + Ok(Self { + config, + http, + api_key, + api_base, + }) + } + + // ==================== 纯函数(协议构造/解析,可单测) ==================== + + /// 构造 /listen 端点完整 URL(`{api_base}/listen`) + fn build_listen_url(api_base: &str) -> String { + format!( + "{base}{path}", + base = api_base.trim_end_matches('/'), + path = LISTEN_PATH, + ) + } + + /// 从文件扩展名推断音频 MIME 类型 + /// + /// 对应 Python v1 `get_audio_bytes` 的 `mime_map`。未知扩展名回退 `audio/wav`。 + fn mime_from_extension(ext: &str) -> &'static str { + match ext.to_lowercase().as_str() { + "mp3" => "audio/mpeg", + "wav" => "audio/wav", + "ogg" => "audio/ogg", + "flac" => "audio/flac", + "m4a" | "mp4" => "audio/mp4", + "webm" => "audio/webm", + "aac" => "audio/aac", + "opus" => "audio/opus", + _ => "audio/wav", + } + } + + /// 从 `extra` 读取布尔参数(支持 bool 或字符串 "true"/"false"),未设置返回 `default` + /// + /// Deepgram 特有参数(smart_format/punctuate/diarize 等)通过 `extra` 透传, + /// 此函数统一处理 bool 与字符串两种形式。 + fn get_bool_extra( + extra: &std::collections::HashMap, + key: &str, + default: bool, + ) -> bool { + match extra.get(key) { + Some(Value::Bool(b)) => *b, + Some(Value::String(s)) => matches!(s.to_lowercase().as_str(), "true" | "1"), + _ => default, + } + } + + /// 构造 Deepgram 查询参数列表 + /// + /// 从 `TranscribeRequest` 提取 Deepgram 特有参数: + /// - `model`(必传,使用已解析的 model,空则由调用方填默认值) + /// - `smart_format`(默认 true) + /// - `punctuate`(默认 true) + /// - `language`(请求显式设置时传) + /// - `diarize`(默认 false,true 时传 "true") + /// - `profanity_filter`(默认 false) + /// - `utterances`(默认 false) + /// - `detect_language`(默认 false) + /// - `keywords`(列表,每个元素作为一个重复的 `keywords` 参数) + /// + /// 与 Python v1 一致:布尔参数仅在为 true 时加入 query(false 省略,依赖服务端默认)。 + fn build_query_params(req: &TranscribeRequest, model: &str) -> Vec<(String, String)> { + let mut params: Vec<(String, String)> = Vec::new(); + params.push(("model".to_string(), model.to_string())); + + if Self::get_bool_extra(&req.extra, "smart_format", true) { + params.push(("smart_format".to_string(), "true".to_string())); + } + if Self::get_bool_extra(&req.extra, "punctuate", true) { + params.push(("punctuate".to_string(), "true".to_string())); + } + if let Some(lang) = &req.language { + if !lang.is_empty() { + params.push(("language".to_string(), lang.clone())); + } + } + if Self::get_bool_extra(&req.extra, "diarize", false) { + params.push(("diarize".to_string(), "true".to_string())); + } + if Self::get_bool_extra(&req.extra, "profanity_filter", false) { + params.push(("profanity_filter".to_string(), "true".to_string())); + } + if Self::get_bool_extra(&req.extra, "utterances", false) { + params.push(("utterances".to_string(), "true".to_string())); + } + if Self::get_bool_extra(&req.extra, "detect_language", false) { + params.push(("detect_language".to_string(), "true".to_string())); + } + // keywords:列表,每个关键词作为一个重复的 query 参数 + if let Some(keywords) = req.extra.get("keywords").and_then(|v| v.as_array()) { + for kw in keywords { + if let Some(s) = kw.as_str() { + params.push(("keywords".to_string(), s.to_string())); + } + } + } + + params + } + + /// 解析 `FileInput` 为 Deepgram 要求的音频输入形式 + /// + /// - `Bytes` → 直接 body(Content-Type 默认 `audio/wav`) + /// - `Base64` → 解码后 body(Content-Type 默认 `audio/wav`) + /// - `Path` → 读取文件后 body(按扩展名推断 Content-Type);文件不存在 → Validation 错误 + /// - `Url` → 用 `url` query param(不传 body) + async fn resolve_audio_input(file: &FileInput) -> Result { + match file { + FileInput::Bytes(b) => Ok(AudioInput::Bytes(b.clone(), "audio/wav".to_string())), + FileInput::Base64(s) => { + let bytes = crate::util::decode_base64(s) + .map_err(|e| AibridgeError::validation(format!("Base64 音频解码失败: {e}")))?; + Ok(AudioInput::Bytes(bytes, "audio/wav".to_string())) + } + FileInput::Path(p) => { + let path = std::path::Path::new(p); + if !path.exists() { + return Err(AibridgeError::validation(format!("音频文件不存在: {p}"))); + } + let bytes = tokio::fs::read(path) + .await + .map_err(|e| AibridgeError::validation(format!("读取音频文件失败: {e}")))?; + let ext = path.extension().and_then(|e| e.to_str()).unwrap_or(""); + let mime = Self::mime_from_extension(ext).to_string(); + Ok(AudioInput::Bytes(bytes, mime)) + } + FileInput::Url(u) => Ok(AudioInput::RemoteUrl(u.clone())), + } + } + + /// 解析 Deepgram 响应为统一 `TranscriptionResult` + /// + /// Deepgram 响应结构: + /// ```json + /// { + /// "results": { + /// "channels": [{ + /// "alternatives": [{ + /// "transcript": "...", + /// "confidence": 0.99, + /// "words": [{ "word", "start", "end", "confidence" }], + /// "paragraphs": { "paragraphs": [{ "speaker", "sentences": [{ "text", "start", "end" }] }] } + /// }] + /// }] + /// }, + /// "metadata": { "duration": 3.5, "model_info": { "name": "nova-2-general", "language": "en" } } + /// } + /// ``` + /// + /// - `text`:拼接所有 channel 的 `alternatives[0].transcript`(空格分隔) + /// - `language`:取自 `metadata.model_info.language` + /// - `duration`:取自 `metadata.duration` + /// - `words`:合并所有 channel 的 `alternatives[0].words` + /// - `segments`:从 `alternatives[0].paragraphs.paragraphs[].sentences[]` 提取, + /// 带说话人信息(`paragraph.speaker`) + fn parse_deepgram_response(result: &Value, model: &str) -> TranscriptionResult { + let results = result.get("results").unwrap_or(&Value::Null); + let channels = results.get("channels").and_then(|c| c.as_array()); + + let mut full_text_parts: Vec = Vec::new(); + let mut all_words: Vec = Vec::new(); + let mut all_segments: Vec = Vec::new(); + + let mut language: Option = None; + let mut duration: Option = None; + + // metadata:duration + model_info.language + if let Some(metadata) = result.get("metadata") { + duration = metadata.get("duration").and_then(|d| d.as_f64()); + if let Some(model_info) = metadata.get("model_info") { + language = model_info + .get("language") + .and_then(|l| l.as_str()) + .map(str::to_owned); + } + } + + if let Some(channels) = channels { + for channel in channels { + let Some(alts) = channel.get("alternatives").and_then(|a| a.as_array()) else { + continue; + }; + let Some(alt) = alts.first() else { + continue; + }; + + // transcript + if let Some(transcript) = alt.get("transcript").and_then(|t| t.as_str()) { + if !transcript.is_empty() { + full_text_parts.push(transcript.to_string()); + } + } + + // words + if let Some(words) = alt.get("words").and_then(|w| w.as_array()) { + for w in words { + let word = w + .get("word") + .and_then(|x| x.as_str()) + .unwrap_or("") + .to_string(); + let start = w.get("start").and_then(|x| x.as_f64()).unwrap_or(0.0); + let end = w.get("end").and_then(|x| x.as_f64()).unwrap_or(0.0); + let confidence = w.get("confidence").and_then(|x| x.as_f64()); + all_words.push(TranscriptionWord { + word, + start, + end, + confidence, + }); + } + } + + // paragraphs → sentences → segments + if let Some(paras) = alt + .get("paragraphs") + .and_then(|p| p.get("paragraphs")) + .and_then(|p| p.as_array()) + { + for para in paras { + let speaker = para + .get("speaker") + .and_then(|s| s.as_i64()) + .map(|s| s.to_string()); + if let Some(sentences) = para.get("sentences").and_then(|s| s.as_array()) { + for sent in sentences { + let id = all_segments.len() as u32; + let start = + sent.get("start").and_then(|x| x.as_f64()).unwrap_or(0.0); + let end = sent.get("end").and_then(|x| x.as_f64()).unwrap_or(0.0); + let text = sent + .get("text") + .and_then(|x| x.as_str()) + .unwrap_or("") + .to_string(); + all_segments.push(TranscriptionSegment { + id, + start, + end, + text, + confidence: None, + speaker: speaker.clone(), + }); + } + } + } + } + } + } + + let full_text = full_text_parts.join(" "); + TranscriptionResult { + text: full_text, + language, + duration, + segments: if all_segments.is_empty() { + None + } else { + Some(all_segments) + }, + words: if all_words.is_empty() { + None + } else { + Some(all_words) + }, + task: "transcribe".to_string(), + usage: None, + model: Some(model.to_string()), + } + } + + /// 将 Deepgram HTTP 错误响应映射为 `AibridgeError` + /// + /// 映射规则见模块文档"错误映射"小节。 + fn map_deepgram_error(status: u16, body: &str) -> AibridgeError { + match status { + 401 => AibridgeError::authentication(format!("Deepgram API key (Token) 无效: {body}")), + 403 => AibridgeError::authentication(format!( + "Deepgram API key 权限不足或账户已停用: {body}" + )), + 429 => AibridgeError::rate_limit(format!("Deepgram 限流或配额耗尽: {body}")), + 400 | 422 => AibridgeError::validation(format!("Deepgram 请求参数错误: {body}")), + s if s >= 500 => AibridgeError::api(s, format!("Deepgram 服务错误 ({s}): {body}")), + s => AibridgeError::api(s, format!("Deepgram HTTP {s}: {body}")), + } + } + + /// 获取 API Key,为空时返 Validation 错误 + fn require_api_key(&self) -> Result<&str> { + self.api_key + .as_deref() + .filter(|k| !k.trim().is_empty()) + .ok_or_else(|| AibridgeError::validation("Deepgram 需要 API key(Token)")) + } +} + +#[async_trait] +impl Adapter for DeepgramAdapter { + fn provider_type(&self) -> &str { + PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::AudioTranscribe); + caps + } + + /// 需 API Key 认证 + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new 时已建好,保持 no-op + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // 无长连接资源需释放 + Ok(()) + } + + /// 语音转文字(Deepgram ASR) + /// + /// - `translate=true` 时返回 Validation 错误(Deepgram 不支持翻译) + /// - 二进制音频直接作为 body 上传;URL 输入用 `url` query param + async fn transcribe(&self, req: TranscribeRequest) -> Result { + // Deepgram 不支持翻译为英文(Python v1 也未实现 translate) + if req.translate { + return Err(AibridgeError::validation( + "Deepgram 不支持翻译为英文(translate=true),仅支持语音转文字(transcribe)", + )); + } + + let api_key = self.require_api_key()?; + + // 模型缺省回退到默认模型 + let model = if req.model.is_empty() { + DEFAULT_MODEL.to_string() + } else { + req.model.clone() + }; + + // 解析音频输入(可能读文件/解码 base64,异步) + let audio_input = Self::resolve_audio_input(&req.file).await?; + + // 构造查询参数(model 用已解析的非空值) + let mut params = Self::build_query_params(&req, &model); + // URL 输入:追加 url query param,不传 body + if let AudioInput::RemoteUrl(u) = &audio_input { + params.push(("url".to_string(), u.clone())); + } + + let url = Self::build_listen_url(&self.api_base); + + // 构造请求:Authorization: Token ,二进制 body(若有) + let mut request = self + .http + .inner() + .post(&url) + .header("Authorization", format!("Token {api_key}")); + + if let AudioInput::Bytes(bytes, content_type) = &audio_input { + request = request + .header("Content-Type", content_type.as_str()) + .body(bytes.clone()); + } + request = request.query(¶ms); + + let resp = request.send().await.map_err(map_reqwest_err)?; + let status = resp.status(); + if !status.is_success() { + let body = resp.text().await.unwrap_or_default(); + return Err(Self::map_deepgram_error(status.as_u16(), &body)); + } + + let body_bytes = resp.bytes().await.map_err(map_reqwest_err)?; + let result: Value = serde_json::from_slice(&body_bytes) + .map_err(|e| AibridgeError::validation(format!("Deepgram 响应解析失败: {e}")))?; + + Ok(Self::parse_deepgram_response(&result, &model)) + } + + /// 列出 Deepgram ASR 模型(硬编码列表,无 /models 拉取需求) + async fn list_models(&self, filter: Option) -> Result> { + let models = deepgram_models(); + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 辅助函数 ==================== + +/// 将 reqwest::Error 映射为 AibridgeError(超时 → Timeout,其余 → Network) +fn map_reqwest_err(e: reqwest::Error) -> AibridgeError { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } +} + +/// Deepgram ASR 模型硬编码列表 +/// +/// 对应 Python v1 `DeepgramAdapter.list_models`。含 Nova-3/Nova-2 全系列 + +/// Whisper 托管版 + 旧版 Enhanced/Base。 +fn deepgram_models() -> Vec { + fn m(id: &str, name: &str, desc: &str) -> ModelInfo { + ModelInfo { + id: id.into(), + name: name.into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_transcribe".into()], + max_tokens: None, + supports_streaming: false, + description: Some(desc.into()), + created: None, + } + } + vec![ + m("nova-3", "Nova 3", "Nova 3 最新通用模型,最高准确率"), + m("nova-2", "Nova 2", "Nova 2 通用模型,推荐默认使用"), + m("nova-2-general", "Nova 2 General", "Nova 2 通用场景"), + m( + "nova-2-meeting", + "Nova 2 Meeting", + "Nova 2 会议场景优化(多人对话)", + ), + m( + "nova-2-phonecall", + "Nova 2 Phone Call", + "Nova 2 电话通话优化(8kHz 音频)", + ), + m( + "nova-2-conversationalai", + "Nova 2 Conversational AI", + "Nova 2 对话 AI/语音助手优化", + ), + m( + "nova-2-video", + "Nova 2 Video", + "Nova 2 视频/播客/多说话人场景", + ), + m("nova-2-medical", "Nova 2 Medical", "Nova 2 医疗领域优化"), + m("nova-2-finance", "Nova 2 Finance", "Nova 2 金融领域优化"), + m( + "nova-2-drivethru", + "Nova 2 Drive-Thru", + "Nova 2 餐厅免下车窗口优化", + ), + m( + "whisper-large", + "Whisper Large (Deepgram)", + "OpenAI Whisper Large 托管版", + ), + m( + "whisper-medium", + "Whisper Medium (Deepgram)", + "OpenAI Whisper Medium 托管版", + ), + m( + "whisper-small", + "Whisper Small (Deepgram)", + "OpenAI Whisper Small 托管版", + ), + m( + "enhanced", + "Enhanced", + "Enhanced 增强模型(旧版,兼容使用)", + ), + m("base", "Base", "Base 基础模型(最快、成本最低)"), + ] +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::config::ClientOptions; + use crate::model::audio::SpeechRequest; + use crate::model::chat::ChatRequest; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + use std::collections::HashMap; + + /// 构造测试用适配器(base_url 指向 mockito server,api_key 可选) + fn make_adapter(base_url: Option, api_key: Option<&str>) -> DeepgramAdapter { + let mut opts = ClientOptions::builder(); + if let Some(u) = base_url { + opts = opts.base_url(u); + } + if let Some(k) = api_key { + opts = opts.api_key(k); + } + let config = ProviderConfig::from_options(PROVIDER_TYPE, opts.build()); + DeepgramAdapter::new(config).expect("构造 DeepgramAdapter 失败") + } + + /// 构造一个最小 Deepgram 成功响应 JSON + fn sample_response() -> Value { + serde_json::json!({ + "results": { + "channels": [{ + "alternatives": [{ + "transcript": "hello world", + "confidence": 0.99, + "words": [ + {"word": "hello", "start": 0.0, "end": 0.5, "confidence": 0.98}, + {"word": "world", "start": 0.5, "end": 1.0, "confidence": 0.97} + ] + }] + }] + }, + "metadata": { + "duration": 1.0, + "model_info": {"name": "nova-2-general", "language": "en"} + } + }) + } + + // ============ 基本属性 ============ + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter(None, Some("test-key")); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_transcribe_only() { + let adapter = make_adapter(None, Some("test-key")); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::AudioTranscribe)); + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::AudioSpeech)); + assert!(!caps.contains(&Capabilities::ListVoices)); + } + + #[test] + fn provider_type_and_name() { + let adapter = make_adapter(None, Some("test-key")); + assert_eq!(adapter.provider_type(), "deepgram"); + assert_eq!(adapter.provider_name(), "Deepgram"); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter(None, Some("test-key")); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ list_models ============ + + #[tokio::test] + async fn list_models_returns_deepgram_models() { + let adapter = make_adapter(None, Some("test-key")); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 15); + assert_eq!(models[0].id, "nova-3"); + assert_eq!(models[1].id, "nova-2"); + assert_eq!(models[0].model_type, ModelType::Audio); + assert_eq!(models[0].provider, "deepgram"); + } + + #[tokio::test] + async fn list_models_filter_by_audio() { + let adapter = make_adapter(None, Some("test-key")); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 15); + let chat = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat.is_empty()); + } + + // ============ 不支持能力(默认实现) ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = ChatRequest::builder("m", vec![]).build(); + assert!(matches!( + adapter.chat(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn image_generate_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = ImageRequest::builder("m", "p").build(); + assert!(matches!( + adapter.image_generate(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = VideoRequest::builder("m", "p").build(); + assert!(matches!( + adapter.video_create(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn embed_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + assert!(matches!( + adapter.embed(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn speech_returns_unsupported() { + let adapter = make_adapter(None, Some("test-key")); + let req = SpeechRequest::builder("aura", "hi", "voice").build(); + assert!(matches!( + adapter.speech(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + // ============ build_listen_url ============ + + #[test] + fn build_listen_url_is_correct() { + let url = DeepgramAdapter::build_listen_url("https://api.deepgram.com/v1"); + assert_eq!(url, "https://api.deepgram.com/v1/listen"); + } + + #[test] + fn build_listen_url_strips_trailing_slash() { + let url = DeepgramAdapter::build_listen_url("https://api.deepgram.com/v1/"); + assert_eq!(url, "https://api.deepgram.com/v1/listen"); + } + + // ============ mime_from_extension ============ + + #[test] + fn mime_from_extension_known() { + assert_eq!(DeepgramAdapter::mime_from_extension("mp3"), "audio/mpeg"); + assert_eq!(DeepgramAdapter::mime_from_extension("wav"), "audio/wav"); + assert_eq!(DeepgramAdapter::mime_from_extension("ogg"), "audio/ogg"); + assert_eq!(DeepgramAdapter::mime_from_extension("flac"), "audio/flac"); + assert_eq!(DeepgramAdapter::mime_from_extension("m4a"), "audio/mp4"); + assert_eq!(DeepgramAdapter::mime_from_extension("mp4"), "audio/mp4"); + assert_eq!(DeepgramAdapter::mime_from_extension("webm"), "audio/webm"); + assert_eq!(DeepgramAdapter::mime_from_extension("aac"), "audio/aac"); + assert_eq!(DeepgramAdapter::mime_from_extension("opus"), "audio/opus"); + } + + #[test] + fn mime_from_extension_unknown_defaults_wav() { + assert_eq!(DeepgramAdapter::mime_from_extension("xyz"), "audio/wav"); + assert_eq!(DeepgramAdapter::mime_from_extension(""), "audio/wav"); + } + + #[test] + fn mime_from_extension_case_insensitive() { + assert_eq!(DeepgramAdapter::mime_from_extension("MP3"), "audio/mpeg"); + assert_eq!(DeepgramAdapter::mime_from_extension("WAV"), "audio/wav"); + } + + // ============ get_bool_extra ============ + + #[test] + fn get_bool_extra_default_when_absent() { + let extra = HashMap::new(); + assert!(DeepgramAdapter::get_bool_extra( + &extra, + "smart_format", + true + )); + assert!(!DeepgramAdapter::get_bool_extra(&extra, "diarize", false)); + } + + #[test] + fn get_bool_extra_reads_bool() { + let mut extra = HashMap::new(); + extra.insert("smart_format".to_string(), serde_json::json!(false)); + assert!(!DeepgramAdapter::get_bool_extra( + &extra, + "smart_format", + true + )); + } + + #[test] + fn get_bool_extra_reads_string_true() { + let mut extra = HashMap::new(); + extra.insert("diarize".to_string(), serde_json::json!("true")); + assert!(DeepgramAdapter::get_bool_extra(&extra, "diarize", false)); + extra.insert("diarize".to_string(), serde_json::json!("TRUE")); + assert!(DeepgramAdapter::get_bool_extra(&extra, "diarize", false)); + } + + #[test] + fn get_bool_extra_reads_string_false() { + let mut extra = HashMap::new(); + extra.insert("smart_format".to_string(), serde_json::json!("false")); + assert!(!DeepgramAdapter::get_bool_extra( + &extra, + "smart_format", + true + )); + } + + // ============ build_query_params ============ + + #[test] + fn build_query_params_defaults_smart_format_and_punctuate() { + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + // model 总是存在 + assert!(params.contains(&("model".to_string(), "nova-2".to_string()))); + // smart_format / punctuate 默认 true + assert!(params.contains(&("smart_format".to_string(), "true".to_string()))); + assert!(params.contains(&("punctuate".to_string(), "true".to_string()))); + // 默认 false 的不出现 + assert!(!params.iter().any(|(k, _)| k == "diarize")); + assert!(!params.iter().any(|(k, _)| k == "profanity_filter")); + assert!(!params.iter().any(|(k, _)| k == "utterances")); + assert!(!params.iter().any(|(k, _)| k == "detect_language")); + assert!(!params.iter().any(|(k, _)| k == "language")); + } + + #[test] + fn build_query_params_disables_smart_format() { + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1])) + .extra("smart_format", false) + .extra("punctuate", false) + .build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + assert!(!params.iter().any(|(k, _)| k == "smart_format")); + assert!(!params.iter().any(|(k, _)| k == "punctuate")); + } + + #[test] + fn build_query_params_includes_language() { + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1])) + .language("zh") + .build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + assert!(params.contains(&("language".to_string(), "zh".to_string()))); + } + + #[test] + fn build_query_params_enables_optional_flags() { + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1])) + .extra("diarize", true) + .extra("profanity_filter", true) + .extra("utterances", true) + .extra("detect_language", true) + .build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + assert!(params.contains(&("diarize".to_string(), "true".to_string()))); + assert!(params.contains(&("profanity_filter".to_string(), "true".to_string()))); + assert!(params.contains(&("utterances".to_string(), "true".to_string()))); + assert!(params.contains(&("detect_language".to_string(), "true".to_string()))); + } + + #[test] + fn build_query_params_keywords_repeated() { + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1])) + .extra("keywords", serde_json::json!(["foo:2", "bar:1"])) + .build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + // keywords 出现两次(重复 query 参数) + let kw_values: Vec = params + .iter() + .filter(|(k, _)| k == "keywords") + .map(|(_, v)| v.clone()) + .collect(); + assert_eq!(kw_values.len(), 2); + assert!(kw_values.contains(&"foo:2".to_string())); + assert!(kw_values.contains(&"bar:1".to_string())); + } + + #[test] + fn build_query_params_uses_resolved_model() { + // 空 model 的请求,调用方传入解析后的默认模型 + let req = TranscribeRequest::builder("", FileInput::bytes(vec![1])).build(); + let params = DeepgramAdapter::build_query_params(&req, DEFAULT_MODEL); + assert!(params.contains(&("model".to_string(), "nova-2".to_string()))); + } + + // ============ resolve_audio_input ============ + + #[tokio::test] + async fn resolve_audio_input_bytes() { + let input = DeepgramAdapter::resolve_audio_input(&FileInput::bytes(vec![1, 2, 3])) + .await + .unwrap(); + match input { + AudioInput::Bytes(b, ct) => { + assert_eq!(b, vec![1, 2, 3]); + assert_eq!(ct, "audio/wav"); + } + _ => panic!("应为 Bytes 变体"), + } + } + + #[tokio::test] + async fn resolve_audio_input_base64() { + let encoded = crate::util::encode_base64(b"audio-bytes"); + let input = DeepgramAdapter::resolve_audio_input(&FileInput::base64(encoded)) + .await + .unwrap(); + match input { + AudioInput::Bytes(b, _) => assert_eq!(b, b"audio-bytes".to_vec()), + _ => panic!("应为 Bytes 变体"), + } + } + + #[tokio::test] + async fn resolve_audio_input_base64_invalid_returns_validation() { + let err = DeepgramAdapter::resolve_audio_input(&FileInput::base64("!!!not-base64!!!")) + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn resolve_audio_input_url() { + let input = + DeepgramAdapter::resolve_audio_input(&FileInput::url("https://example.com/audio.mp3")) + .await + .unwrap(); + match input { + AudioInput::RemoteUrl(u) => assert_eq!(u, "https://example.com/audio.mp3"), + _ => panic!("应为 RemoteUrl 变体"), + } + } + + #[tokio::test] + async fn resolve_audio_input_path_reads_file() { + // 写一个临时文件再读取 + let path = std::env::temp_dir().join("aibridge_deepgram_test_read.wav"); + std::fs::write(&path, b"wav-bytes").unwrap(); + let input = DeepgramAdapter::resolve_audio_input(&FileInput::path(path.to_str().unwrap())) + .await + .unwrap(); + match input { + AudioInput::Bytes(b, ct) => { + assert_eq!(b, b"wav-bytes".to_vec()); + assert_eq!(ct, "audio/wav"); + } + _ => panic!("应为 Bytes 变体"), + } + let _ = std::fs::remove_file(&path); + } + + #[tokio::test] + async fn resolve_audio_input_path_mp3_mime() { + let path = std::env::temp_dir().join("aibridge_deepgram_test_read.mp3"); + std::fs::write(&path, b"mp3-bytes").unwrap(); + let input = DeepgramAdapter::resolve_audio_input(&FileInput::path(path.to_str().unwrap())) + .await + .unwrap(); + match input { + AudioInput::Bytes(_, ct) => assert_eq!(ct, "audio/mpeg"), + _ => panic!("应为 Bytes 变体"), + } + let _ = std::fs::remove_file(&path); + } + + #[tokio::test] + async fn resolve_audio_input_path_nonexistent_returns_validation() { + let err = DeepgramAdapter::resolve_audio_input(&FileInput::path( + "/tmp/aibridge_nonexistent_file_xyz.wav", + )) + .await + .unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + // ============ parse_deepgram_response ============ + + #[test] + fn parse_response_basic_with_words() { + let result = DeepgramAdapter::parse_deepgram_response(&sample_response(), "nova-2"); + assert_eq!(result.text, "hello world"); + assert_eq!(result.language.as_deref(), Some("en")); + assert!((result.duration.unwrap() - 1.0).abs() < f64::EPSILON); + assert_eq!(result.task, "transcribe"); + assert_eq!(result.model.as_deref(), Some("nova-2")); + let words = result.words.expect("应有 words"); + assert_eq!(words.len(), 2); + assert_eq!(words[0].word, "hello"); + assert!((words[0].start - 0.0).abs() < f64::EPSILON); + assert!((words[0].end - 0.5).abs() < f64::EPSILON); + assert!((words[0].confidence.unwrap() - 0.98).abs() < f64::EPSILON); + assert!(result.segments.is_none()); + } + + #[test] + fn parse_response_with_segments_and_speaker() { + let json = serde_json::json!({ + "results": { + "channels": [{ + "alternatives": [{ + "transcript": "hello world", + "paragraphs": { + "paragraphs": [{ + "speaker": 0, + "sentences": [ + {"text": "hello", "start": 0.0, "end": 0.5}, + {"text": "world", "start": 0.5, "end": 1.0} + ] + }] + } + }] + }] + }, + "metadata": {"duration": 1.0, "model_info": {"language": "en"}} + }); + let result = DeepgramAdapter::parse_deepgram_response(&json, "nova-3"); + assert_eq!(result.text, "hello world"); + let segments = result.segments.expect("应有 segments"); + assert_eq!(segments.len(), 2); + assert_eq!(segments[0].id, 0); + assert_eq!(segments[0].text, "hello"); + assert!((segments[0].start - 0.0).abs() < f64::EPSILON); + assert!((segments[0].end - 0.5).abs() < f64::EPSILON); + assert_eq!(segments[1].id, 1); + assert_eq!(segments[1].text, "world"); + assert_eq!(segments[0].speaker.as_deref(), Some("0")); + assert_eq!(segments[1].speaker.as_deref(), Some("0")); + } + + #[test] + fn parse_response_empty_channels() { + let json = serde_json::json!({ + "results": {"channels": []}, + "metadata": {"duration": 0.0} + }); + let result = DeepgramAdapter::parse_deepgram_response(&json, "nova-2"); + assert!(result.text.is_empty()); + assert!(result.words.is_none()); + assert!(result.segments.is_none()); + assert!(result.language.is_none()); + assert!((result.duration.unwrap() - 0.0).abs() < f64::EPSILON); + } + + #[test] + fn parse_response_multiple_channels_concatenates() { + let json = serde_json::json!({ + "results": { + "channels": [ + {"alternatives": [{"transcript": "hello"}]}, + {"alternatives": [{"transcript": "world"}]} + ] + }, + "metadata": {} + }); + let result = DeepgramAdapter::parse_deepgram_response(&json, "nova-2"); + assert_eq!(result.text, "hello world"); + } + + #[test] + fn parse_response_missing_metadata() { + let json = serde_json::json!({ + "results": { + "channels": [{"alternatives": [{"transcript": "hi"}]}] + } + }); + let result = DeepgramAdapter::parse_deepgram_response(&json, "nova-2"); + assert_eq!(result.text, "hi"); + assert!(result.duration.is_none()); + assert!(result.language.is_none()); + } + + #[test] + fn parse_response_empty_transcript_skipped() { + let json = serde_json::json!({ + "results": { + "channels": [{"alternatives": [{"transcript": ""}]}] + }, + "metadata": {} + }); + let result = DeepgramAdapter::parse_deepgram_response(&json, "nova-2"); + assert!(result.text.is_empty()); + } + + // ============ map_deepgram_error ============ + + #[test] + fn map_error_401_is_authentication() { + let err = DeepgramAdapter::map_deepgram_error(401, "invalid token"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_error_403_is_authentication() { + let err = DeepgramAdapter::map_deepgram_error(403, "forbidden"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_error_429_is_rate_limit() { + let err = DeepgramAdapter::map_deepgram_error(429, "slow down"); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[test] + fn map_error_400_is_validation() { + let err = DeepgramAdapter::map_deepgram_error(400, "bad request"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_error_422_is_validation() { + let err = DeepgramAdapter::map_deepgram_error(422, "unprocessable"); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[test] + fn map_error_500_is_api() { + let err = DeepgramAdapter::map_deepgram_error(500, "server error"); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[test] + fn map_error_503_is_api() { + let err = DeepgramAdapter::map_deepgram_error(503, "unavailable"); + assert!(matches!(err, AibridgeError::Api { status: 503, .. })); + } + + #[test] + fn map_error_other_4xx_is_api() { + let err = DeepgramAdapter::map_deepgram_error(418, "teapot"); + assert!(matches!(err, AibridgeError::Api { status: 418, .. })); + } + + // ============ transcribe(mockito) ============ + + #[tokio::test] + async fn transcribe_bytes_success() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", LISTEN_PATH) + .match_header("authorization", "Token test-key") + .match_header("content-type", "audio/wav") + .match_query(mockito::Matcher::AllOf(vec![ + mockito::Matcher::UrlEncoded("model".into(), "nova-2".into()), + mockito::Matcher::UrlEncoded("smart_format".into(), "true".into()), + ])) + .match_body(mockito::Matcher::from(b"audio-bytes".to_vec())) + .with_status(200) + .with_header("content-type", "application/json") + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = + TranscribeRequest::builder("nova-2", FileInput::bytes(b"audio-bytes".to_vec())).build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "hello world"); + assert_eq!(result.language.as_deref(), Some("en")); + assert_eq!(result.model.as_deref(), Some("nova-2")); + assert_eq!(result.task, "transcribe"); + assert!(result.words.is_some()); + mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_url_uses_url_query_param_no_body() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", LISTEN_PATH) + .match_header("authorization", "Token test-key") + .match_query(mockito::Matcher::AllOf(vec![ + mockito::Matcher::UrlEncoded("url".into(), "https://example.com/audio.mp3".into()), + mockito::Matcher::UrlEncoded("model".into(), "nova-2".into()), + ])) + .match_body(mockito::Matcher::Missing) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = + TranscribeRequest::builder("nova-2", FileInput::url("https://example.com/audio.mp3")) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "hello world"); + mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_with_language_query_param() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::UrlEncoded("language".into(), "zh".into())) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])) + .language("zh") + .build(); + let _ = adapter.transcribe(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_smart_format_disabled_not_in_query() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::UrlEncoded( + "model".into(), + "nova-2".into(), + )) + // smart_format 不应出现(mock 默认允许其他 query,这里仅断言 model 存在) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])) + .extra("smart_format", false) + .build(); + let params = DeepgramAdapter::build_query_params(&req, "nova-2"); + assert!(!params.iter().any(|(k, _)| k == "smart_format")); + let _ = adapter.transcribe(req).await.unwrap(); + mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_uses_default_model_when_empty() { + let mut server = mockito::Server::new_async().await; + let mock = server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::UrlEncoded( + "model".into(), + DEFAULT_MODEL.into(), + )) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("", FileInput::bytes(vec![1, 2, 3])).build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.model.as_deref(), Some(DEFAULT_MODEL)); + mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_path_input_uploads_file_bytes() { + let mut server = mockito::Server::new_async().await; + // 写临时 wav 文件 + let path = std::env::temp_dir().join("aibridge_deepgram_transcribe_test.wav"); + std::fs::write(&path, b"wav-file-content").unwrap(); + let mock = server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .match_header("content-type", "audio/wav") + .match_body(mockito::Matcher::from(b"wav-file-content".to_vec())) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = + TranscribeRequest::builder("nova-2", FileInput::path(path.to_str().unwrap())).build(); + let _ = adapter.transcribe(req).await.unwrap(); + mock.assert_async().await; + let _ = std::fs::remove_file(&path); + } + + #[tokio::test] + async fn transcribe_error_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(401) + .with_body("{\"err_msg\":\"invalid api key\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("bad-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn transcribe_error_429_returns_rate_limit() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(429) + .with_body("{\"err_msg\":\"quota exceeded\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn transcribe_error_400_returns_validation() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(400) + .with_body("{\"err_msg\":\"invalid model\"}") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("bad-model", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn transcribe_error_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(500) + .with_body("internal error") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[tokio::test] + async fn transcribe_error_403_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(403) + .with_body("forbidden") + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + // ============ translate(Deepgram 不支持) ============ + + #[tokio::test] + async fn transcribe_translate_true_returns_validation() { + let adapter = make_adapter(None, Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])) + .translate(true) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn transcribe_translate_false_proceeds() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", LISTEN_PATH) + .match_query(mockito::Matcher::Any) + .with_status(200) + .with_body(sample_response().to_string()) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url()), Some("test-key")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])) + .translate(false) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "hello world"); + } + + // ============ API key 缺失 ============ + + #[tokio::test] + async fn transcribe_without_api_key_returns_validation() { + let adapter = make_adapter(Some("https://example.com".into()), None); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn transcribe_with_empty_api_key_returns_validation() { + let adapter = make_adapter(Some("https://example.com".into()), Some("")); + let req = TranscribeRequest::builder("nova-2", FileInput::bytes(vec![1, 2, 3])).build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } +} From 3fb46bed306bf92435a69d04efeb251793e6bbbb Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 02:01:33 +0800 Subject: [PATCH 46/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20assemblyai=20=E9=80=82=E9=85=8D=E5=99=A8?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- .../aibridge-core/src/adapters/assemblyai.rs | 1567 +++++++++++++++++ 1 file changed, 1567 insertions(+) create mode 100644 crates/aibridge-core/src/adapters/assemblyai.rs diff --git a/crates/aibridge-core/src/adapters/assemblyai.rs b/crates/aibridge-core/src/adapters/assemblyai.rs new file mode 100644 index 0000000..58ebdb9 --- /dev/null +++ b/crates/aibridge-core/src/adapters/assemblyai.rs @@ -0,0 +1,1567 @@ +//! AssemblyAI 适配器(企业级 ASR,speaker diarization,独立协议) +//! +//! 对应 Python v1 (agn-sdk) 的 `agn/adapters/audio_adapters.py` 的 `AssemblyAIAdapter`。 +//! +//! AssemblyAI 是企业级语音识别服务,支持说话人分离(speaker diarization)、 +//! 情感分析、PII 脱敏、章节检测、实体识别等丰富语音理解能力。 +//! +//! ## 协议(异步三段式) +//! +//! AssemblyAI 采用异步任务模式,转写流程分三步: +//! 1. 上传音频:`POST /upload`(原始字节,`Content-Type: application/octet-stream`), +//! 返回 `{"upload_url": "..."}` +//! 2. 创建转写任务:`POST /transcript`(JSON body,含 audio_url + 参数), +//! 返回 `{"id": "..."}` +//! 3. 轮询结果:`GET /transcript/{id}`,返回 `{"status": "completed|error|processing|queued", ...}` +//! +//! - 认证:`Authorization: ` header(注意:不是 Bearer,直接传 key) +//! - 文档:https://www.assemblyai.com/docs +//! - 若调用方直接提供 `extra.audio_url`(或 `FileInput::Url`),跳过上传步骤 +//! +//! ## 特性(v1.3.3 保留) +//! +//! - `requires_api_key = true`(需 AssemblyAI API Key) +//! - `speech_model`:best(默认,最高准确率)/ nano(轻量低成本) +//! - 高级参数:speaker_labels / sentiment_analysis / auto_chapters / entity_detection / +//! redact_pii / word_boost / filter_profanity 等,通过 `extra` 透传 +//! - 时间戳:AssemblyAI 的 start/end 为毫秒,归一化为秒;audio_duration 已是秒 + +use std::time::Duration; + +use async_trait::async_trait; +use serde_json::{json, Value}; +use tokio::time::sleep; + +use crate::adapter::{Adapter, Capabilities, CapabilitySet}; +use crate::config::{ClientOptions, ProviderConfig}; +use crate::error::{AibridgeError, Result}; +use crate::http::HttpClient; +use crate::model::audio::{ + TranscribeRequest, TranscriptionResult, TranscriptionSegment, TranscriptionWord, +}; +use crate::model::common::{ModelInfo, ModelType}; +use crate::model::image::FileInput; + +// ==================== 常量 ==================== + +/// Provider 类型标识 +const PROVIDER_TYPE: &str = "assemblyai"; +/// Provider 显示名称 +const PROVIDER_NAME: &str = "AssemblyAI"; +/// 默认 API 基地址(AssemblyAI v2 API) +const DEFAULT_API_BASE: &str = "https://api.assemblyai.com/v2"; +/// 上传音频端点路径 +const UPLOAD_PATH: &str = "/upload"; +/// 创建/查询转写任务端点路径 +const TRANSCRIPT_PATH: &str = "/transcript"; +/// 默认轮询间隔(秒) +const DEFAULT_POLL_INTERVAL: f64 = 1.0; +/// 默认最大轮询次数(对应 Python v1 `MAX_POLLS = 300`) +const DEFAULT_MAX_POLLS: u64 = 300; + +// ==================== AssemblyAiAdapter ==================== + +/// AssemblyAI 适配器 +/// +/// 持有 HTTP 客户端、API key 与基地址。transcribe 流程为上传 → 创建 → 轮询, +/// 均为 per-request HTTP 调用(无长连接)。 +pub struct AssemblyAiAdapter { + /// Provider 配置(保留供未来扩展,当前字段已在构造时提取) + #[allow(dead_code)] + config: ProviderConfig, + /// HTTP 客户端(封装 reqwest,含连接池与超时) + http: HttpClient, + /// API 基地址(默认 `https://api.assemblyai.com/v2`,可由 config.base_url 覆盖) + api_base: String, + /// API Key(Authorization header 值) + api_key: String, +} + +impl AssemblyAiAdapter { + /// 创建 AssemblyAI 适配器 + /// + /// - `config.base_url` 为空时用 `DEFAULT_API_BASE` + /// - `config.api_key` 为 None 时用空串(调用时 API 会返 401) + pub fn new(config: ProviderConfig) -> Result { + let api_base = config + .base_url + .clone() + .filter(|u| !u.trim().is_empty()) + .unwrap_or_else(|| DEFAULT_API_BASE.to_string()); + let api_key = config.api_key.clone().unwrap_or_default(); + let opts = ClientOptions::builder().timeout(config.timeout).build(); + let http = HttpClient::new(&opts)?; + Ok(Self { + config, + http, + api_base, + api_key, + }) + } + + // ==================== 纯函数(协议构造/解析,可单测) ==================== + + /// 解析 speech_model:best / nano,其余默认 best + /// + /// 对应 Python v1 `speech_model = model if model in ("best", "nano") else "best"`。 + fn resolve_speech_model(model: &str) -> String { + match model { + "best" | "nano" => model.to_string(), + _ => "best".to_string(), + } + } + + /// 构造 `/transcript` 请求体 + /// + /// 对应 Python v1 `transcript_request` 构造。必填字段 audio_url + speech_model, + /// 默认开启 punctuate / format_text。可选字段(language_code / speaker_labels / + /// filter_profanity / sentiment_analysis / auto_chapters / entity_detection / + /// redact_pii / word_boost)仅在有值时加入。 + fn build_transcript_payload( + req: &TranscribeRequest, + audio_url: &str, + speech_model: &str, + ) -> Value { + let mut payload = json!({ + "audio_url": audio_url, + "speech_model": speech_model, + "punctuate": req.extra.get("punctuate").and_then(|v| v.as_bool()).unwrap_or(true), + "format_text": req.extra.get("format_text").and_then(|v| v.as_bool()).unwrap_or(true), + }); + + // language_code(req.language 优先,回退 extra.language_code) + let lang = req + .language + .as_deref() + .filter(|l| !l.is_empty()) + .or_else(|| req.extra.get("language_code").and_then(|v| v.as_str())); + if let Some(lang) = lang { + payload["language_code"] = json!(lang); + } + + // 布尔开关参数:仅当显式为 true 时加入 + for &(extra_key, payload_key) in &[ + ("speaker_labels", "speaker_labels"), + ("filter_profanity", "filter_profanity"), + ("sentiment_analysis", "sentiment_analysis"), + ("auto_chapters", "auto_chapters"), + ("entity_detection", "entity_detection"), + ] { + if req + .extra + .get(extra_key) + .and_then(|v| v.as_bool()) + .unwrap_or(false) + { + payload[payload_key] = json!(true); + } + } + + // redact_pii:开启时同时传 redact_pii_policies(默认空数组) + if req + .extra + .get("redact_pii") + .and_then(|v| v.as_bool()) + .unwrap_or(false) + { + payload["redact_pii"] = json!(true); + payload["redact_pii_policies"] = req + .extra + .get("redact_pii_policies") + .cloned() + .unwrap_or_else(|| json!([])); + } + + // word_boost:仅当为 JSON 数组时透传 + if let Some(wb) = req.extra.get("word_boost").and_then(|v| v.as_array()) { + payload["word_boost"] = json!(wb); + } + + payload + } + + /// 解析 AssemblyAI 转写结果为统一 `TranscriptionResult` + /// + /// - text / language_code / audio_duration 直接映射(audio_duration 已是秒) + /// - utterances(speaker_labels 开启时返回)→ segments + 词级 words,start/end 毫秒转秒 + /// - 无 utterances 时回退到顶层 words 数组 + /// - task 由调用方传入(transcribe / translate) + fn parse_response(v: &Value, speech_model: &str, task: &str) -> TranscriptionResult { + let text = v + .get("text") + .and_then(|t| t.as_str()) + .unwrap_or("") + .to_string(); + let language = v + .get("language_code") + .and_then(|l| l.as_str()) + .map(|s| s.to_string()); + let duration = v.get("audio_duration").and_then(|d| d.as_f64()); + + let mut segments: Vec = Vec::new(); + let mut words: Vec = Vec::new(); + + if let Some(utterances) = v.get("utterances").and_then(|u| u.as_array()) { + for (idx, utt) in utterances.iter().enumerate() { + segments.push(TranscriptionSegment { + id: idx as u32, + start: ms_to_secs(utt.get("start")), + end: ms_to_secs(utt.get("end")), + text: utt + .get("text") + .and_then(|t| t.as_str()) + .unwrap_or("") + .to_string(), + confidence: utt.get("confidence").and_then(|c| c.as_f64()), + speaker: utt + .get("speaker") + .and_then(|s| s.as_str()) + .map(|s| s.to_string()), + }); + if let Some(utt_words) = utt.get("words").and_then(|w| w.as_array()) { + for w in utt_words { + words.push(TranscriptionWord { + word: w + .get("text") + .and_then(|t| t.as_str()) + .unwrap_or("") + .to_string(), + start: ms_to_secs(w.get("start")), + end: ms_to_secs(w.get("end")), + confidence: w.get("confidence").and_then(|c| c.as_f64()), + }); + } + } + } + } else if let Some(api_words) = v.get("words").and_then(|w| w.as_array()) { + for w in api_words { + words.push(TranscriptionWord { + word: w + .get("text") + .and_then(|t| t.as_str()) + .unwrap_or("") + .to_string(), + start: ms_to_secs(w.get("start")), + end: ms_to_secs(w.get("end")), + confidence: w.get("confidence").and_then(|c| c.as_f64()), + }); + } + } + + TranscriptionResult { + text, + language, + duration, + segments: if segments.is_empty() { + None + } else { + Some(segments) + }, + words: if words.is_empty() { None } else { Some(words) }, + task: task.to_string(), + usage: None, + model: Some(speech_model.to_string()), + } + } + + // ==================== 网络方法 ==================== + + /// 上传音频字节到 AssemblyAI,返回 upload_url + /// + /// 对应 Python v1 `_upload_audio`:POST /upload,原始字节 body。 + async fn upload_audio(&self, data: Vec) -> Result { + let url = format!( + "{base}{path}", + base = self.api_base.trim_end_matches('/'), + path = UPLOAD_PATH + ); + let resp = self + .http + .inner() + .post(&url) + .header("Authorization", &self.api_key) + .header("Content-Type", "application/octet-stream") + .body(data) + .send() + .await + .map_err(map_reqwest_error)?; + + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_assemblyai_error(status_code, &body)); + } + let v: Value = resp.json().await.map_err(AibridgeError::from)?; + v.get("upload_url") + .and_then(|u| u.as_str()) + .map(|s| s.to_string()) + .ok_or_else(|| AibridgeError::api(0, "AssemblyAI upload 响应缺少 upload_url 字段")) + } + + /// 创建转写任务,返回 transcript id + /// + /// 对应 Python v1 `POST /transcript`。 + async fn create_transcript(&self, payload: &Value) -> Result { + let url = format!( + "{base}{path}", + base = self.api_base.trim_end_matches('/'), + path = TRANSCRIPT_PATH + ); + let resp = self + .http + .inner() + .post(&url) + .header("Authorization", &self.api_key) + .json(payload) + .send() + .await + .map_err(map_reqwest_error)?; + + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_assemblyai_error(status_code, &body)); + } + let v: Value = resp.json().await.map_err(AibridgeError::from)?; + v.get("id") + .and_then(|i| i.as_str()) + .map(|s| s.to_string()) + .ok_or_else(|| AibridgeError::api(0, "AssemblyAI transcript 创建响应缺少 id 字段")) + } + + /// 轮询转写任务结果,直到 completed / error 或超过最大轮询次数 + /// + /// 对应 Python v1 轮询循环。status=completed 返回结果 JSON;status=error 映射为 + /// Api 错误;processing/queued 继续 sleep 后重试;超过 max_polls 返 Timeout。 + async fn poll_transcript( + &self, + transcript_id: &str, + poll_interval: f64, + max_polls: u64, + ) -> Result { + let url = format!( + "{base}{path}/{id}", + base = self.api_base.trim_end_matches('/'), + path = TRANSCRIPT_PATH, + id = transcript_id, + ); + // 负数或 NaN 防御:回退到默认间隔 + let delay_secs = if poll_interval.is_finite() && poll_interval > 0.0 { + poll_interval + } else { + DEFAULT_POLL_INTERVAL + }; + let delay = Duration::from_secs_f64(delay_secs); + + for _ in 0..max_polls { + let resp = self + .http + .inner() + .get(&url) + .header("Authorization", &self.api_key) + .send() + .await + .map_err(map_reqwest_error)?; + + let status = resp.status(); + if !status.is_success() { + let status_code = status.as_u16(); + let body = resp.text().await.unwrap_or_default(); + return Err(map_assemblyai_error(status_code, &body)); + } + let v: Value = resp.json().await.map_err(AibridgeError::from)?; + match v.get("status").and_then(|s| s.as_str()) { + Some("completed") => return Ok(v), + Some("error") => { + let err_msg = v + .get("error") + .and_then(|e| e.as_str()) + .unwrap_or("转写失败"); + return Err(AibridgeError::api( + 500, + format!("AssemblyAI 转写错误: {err_msg}"), + )); + } + _ => { /* processing / queued / 其他:继续轮询 */ } + } + sleep(delay).await; + } + Err(AibridgeError::Timeout) + } + + /// 解析音频输入来源,返回最终用于创建任务的 audio_url + /// + /// 优先级:`extra.audio_url` > `FileInput::Url` > 上传字节(Bytes/Base64/Path)。 + async fn resolve_audio_url(&self, req: &TranscribeRequest) -> Result { + // 1. extra.audio_url 直接指定(跳过上传) + if let Some(url) = req.extra.get("audio_url").and_then(|v| v.as_str()) { + if !url.is_empty() { + return Ok(url.to_string()); + } + } + // 2. FileInput::Url 直接使用 + if let FileInput::Url(u) = &req.file { + return Ok(u.clone()); + } + // 3. 字节类输入:读取后上传 + let data = self.read_file_input_bytes(&req.file).await?; + self.upload_audio(data).await + } + + /// 从 `FileInput` 读取音频字节(Bytes/Base64/Path 三种) + /// + /// `FileInput::Url` 不应走到这里(由 `resolve_audio_url` 提前处理)。 + async fn read_file_input_bytes(&self, file: &FileInput) -> Result> { + match file { + FileInput::Bytes(data) => Ok(data.clone()), + FileInput::Base64(s) => { + let decoded = crate::util::decode_base64(s).map_err(|e| { + AibridgeError::validation(format!("AssemblyAI 音频 Base64 解码失败: {e}")) + })?; + Ok(decoded) + } + FileInput::Path(p) => { + let data = tokio::fs::read(p).await.map_err(|e| { + AibridgeError::validation(format!("AssemblyAI 读取音频文件失败 {p}: {e}")) + })?; + Ok(data) + } + FileInput::Url(_) => Err(AibridgeError::validation( + "AssemblyAI 不应对 Url 输入执行上传", + )), + } + } +} + +#[async_trait] +impl Adapter for AssemblyAiAdapter { + fn provider_type(&self) -> &str { + PROVIDER_TYPE + } + + fn provider_name(&self) -> &str { + PROVIDER_NAME + } + + fn capabilities(&self) -> CapabilitySet { + let mut caps = CapabilitySet::new(); + caps.insert(Capabilities::AudioTranscribe); + caps + } + + /// AssemblyAI 需要 API Key(Authorization header) + fn requires_api_key(&self) -> bool { + true + } + + async fn start(&mut self) -> Result<()> { + // HTTP 客户端在 new() 已构造,无惰性资源需初始化 + Ok(()) + } + + async fn close(&mut self) -> Result<()> { + // 无长连接资源需释放 + Ok(()) + } + + /// 语音转文字(AssemblyAI 异步三段式协议) + /// + /// 流程:解析 audio_url → 创建任务 → 轮询结果 → 解析为统一格式。 + /// `req.translate = true` 时 task 标记为 "translate"(AssemblyAI 不直接支持音频翻译, + /// 但统一接口保留该语义,结果文本语言由 language_code 决定)。 + async fn transcribe(&self, req: TranscribeRequest) -> Result { + let speech_model = Self::resolve_speech_model(&req.model); + let audio_url = self.resolve_audio_url(&req).await?; + let payload = Self::build_transcript_payload(&req, &audio_url, &speech_model); + + let transcript_id = self.create_transcript(&payload).await?; + + let poll_interval = req + .extra + .get("polling_interval") + .and_then(|v| v.as_f64()) + .unwrap_or(DEFAULT_POLL_INTERVAL); + let max_polls = req + .extra + .get("max_polls") + .and_then(|v| v.as_u64()) + .unwrap_or(DEFAULT_MAX_POLLS); + + let result = self + .poll_transcript(&transcript_id, poll_interval, max_polls) + .await?; + + let task = if req.translate { + "translate" + } else { + "transcribe" + }; + Ok(Self::parse_response(&result, &speech_model, task)) + } + + /// 列出 AssemblyAI 模型(无标准 /models 端点,硬编码 best / nano) + /// + /// 对应 Python v1 `list_models`。 + async fn list_models(&self, filter: Option) -> Result> { + let models = vec![ + ModelInfo { + id: "best".into(), + name: "Best".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_transcribe".into()], + max_tokens: None, + supports_streaming: false, + description: Some("AssemblyAI 最高准确率模型(默认),支持所有高级功能".into()), + created: None, + }, + ModelInfo { + id: "nano".into(), + name: "Nano".into(), + model_type: ModelType::Audio, + provider: PROVIDER_TYPE.into(), + capabilities: vec!["audio_transcribe".into()], + max_tokens: None, + supports_streaming: false, + description: Some("AssemblyAI 轻量模型,更低成本,适合简单场景".into()), + created: None, + }, + ]; + Ok(match filter { + Some(t) => models.into_iter().filter(|m| m.model_type == t).collect(), + None => models, + }) + } +} + +// ==================== 错误映射 ==================== + +/// 将 AssemblyAI HTTP 错误响应映射为 AibridgeError +/// +/// 对应 Python v1 `AssemblyAIAdapter._handle_error`。映射规则(与任务规格一致): +/// - 401/403 → Authentication(API Key 无效) +/// - 429 → RateLimit(限流或配额耗尽) +/// - 400 → Validation(参数校验错误,携带 details) +/// - 5xx → Api(服务端错误) +/// - 其余 4xx → Api +fn map_assemblyai_error(status: u16, body: &str) -> AibridgeError { + let message = parse_assemblyai_error_message(body, status); + match status { + 401 | 403 => AibridgeError::authentication(format!( + "AssemblyAI 认证失败(API Key 无效或缺失): {message}" + )), + 429 => AibridgeError::rate_limit(format!("AssemblyAI 限流或配额耗尽: {message}")), + 400 => AibridgeError::validation_with_details( + format!("AssemblyAI 参数校验错误: {message}"), + serde_json::json!({"status": status, "response": body}), + ), + s if s >= 500 => AibridgeError::api(s, format!("AssemblyAI 服务端错误 ({s}): {message}")), + s => AibridgeError::api(s, format!("AssemblyAI HTTP {s}: {message}")), + } +} + +/// 从 AssemblyAI 错误响应体解析错误消息 +/// +/// AssemblyAI 错误响应通常为 `{"error": "..."}`。解析失败时截取前 200 字符或回退 `HTTP {status}`。 +fn parse_assemblyai_error_message(body: &str, status: u16) -> String { + if let Ok(v) = serde_json::from_str::(body) { + if let Some(m) = v.get("error").and_then(|e| e.as_str()) { + return m.to_string(); + } + if let Some(m) = v.get("message").and_then(|m| m.as_str()) { + return m.to_string(); + } + } + if body.trim().is_empty() { + format!("HTTP {status}") + } else { + body.chars().take(200).collect() + } +} + +// ==================== 辅助函数 ==================== + +/// 将 reqwest::Error 映射为 AibridgeError(超时 → Timeout,其余 → Network) +fn map_reqwest_error(e: reqwest::Error) -> AibridgeError { + if e.is_timeout() { + AibridgeError::Timeout + } else { + AibridgeError::Network(e) + } +} + +/// 把 AssemblyAI 的毫秒时间戳转为秒 +/// +/// `start` / `end` 字段为毫秒(整数),`audio_duration` 已是秒(不调用本函数)。 +fn ms_to_secs(v: Option<&Value>) -> f64 { + v.and_then(|x| x.as_f64()).unwrap_or(0.0) / 1000.0 +} + +#[cfg(test)] +mod tests { + use super::*; + use crate::model::audio::TranscribeRequest; + use crate::model::chat::ChatRequest; + use crate::model::image::{FileInput, ImageRequest}; + use crate::model::options::{EmbedInput, EmbedRequest}; + use crate::model::video::VideoRequest; + use std::collections::HashMap; + + /// 构造测试用适配器(base_url 指向 mockito server,None 用默认基地址) + fn make_adapter(base_url: Option) -> AssemblyAiAdapter { + let mut opts = ClientOptions::builder().api_key("test-key"); + if let Some(u) = base_url { + opts = opts.base_url(u); + } + let config = ProviderConfig::from_options(PROVIDER_TYPE, opts.build()); + AssemblyAiAdapter::new(config).expect("构造 AssemblyAiAdapter 失败") + } + + // ============ 基本属性 ============ + + #[test] + fn requires_api_key_is_true() { + let adapter = make_adapter(None); + assert!(adapter.requires_api_key()); + } + + #[test] + fn capabilities_contains_only_transcribe() { + let adapter = make_adapter(None); + let caps = adapter.capabilities(); + assert!(caps.contains(&Capabilities::AudioTranscribe)); + assert!(!caps.contains(&Capabilities::AudioSpeech)); + assert!(!caps.contains(&Capabilities::Chat)); + assert!(!caps.contains(&Capabilities::ImageGenerate)); + assert!(!caps.contains(&Capabilities::ListVoices)); + } + + #[test] + fn provider_type_and_name() { + let adapter = make_adapter(None); + assert_eq!(adapter.provider_type(), "assemblyai"); + assert_eq!(adapter.provider_name(), "AssemblyAI"); + } + + #[tokio::test] + async fn list_models_returns_best_and_nano() { + let adapter = make_adapter(None); + let models = adapter.list_models(None).await.unwrap(); + assert_eq!(models.len(), 2); + let ids: Vec<&str> = models.iter().map(|m| m.id.as_str()).collect(); + assert!(ids.contains(&"best")); + assert!(ids.contains(&"nano")); + for m in &models { + assert_eq!(m.model_type, ModelType::Audio); + assert_eq!(m.provider, "assemblyai"); + } + } + + #[tokio::test] + async fn list_models_filter_by_type() { + let adapter = make_adapter(None); + let audio = adapter.list_models(Some(ModelType::Audio)).await.unwrap(); + assert_eq!(audio.len(), 2); + let chat = adapter.list_models(Some(ModelType::Chat)).await.unwrap(); + assert!(chat.is_empty()); + } + + #[tokio::test] + async fn start_and_close_are_noops() { + let mut adapter = make_adapter(None); + assert!(adapter.start().await.is_ok()); + assert!(adapter.close().await.is_ok()); + } + + // ============ 不支持能力(默认实现) ============ + + #[tokio::test] + async fn chat_returns_unsupported() { + let adapter = make_adapter(None); + let req = ChatRequest::builder("m", vec![]).build(); + assert!(matches!( + adapter.chat(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn image_generate_returns_unsupported() { + let adapter = make_adapter(None); + let req = ImageRequest::builder("m", "p").build(); + assert!(matches!( + adapter.image_generate(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn video_create_returns_unsupported() { + let adapter = make_adapter(None); + let req = VideoRequest::builder("m", "p").build(); + assert!(matches!( + adapter.video_create(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn embed_returns_unsupported() { + let adapter = make_adapter(None); + let req = EmbedRequest { + model: "m".into(), + input: EmbedInput::Single("hi".into()), + dimensions: None, + encoding_format: None, + user: None, + extra: HashMap::new(), + }; + assert!(matches!( + adapter.embed(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + #[tokio::test] + async fn speech_returns_unsupported() { + let adapter = make_adapter(None); + let req = crate::model::audio::SpeechRequest::builder("m", "hi", "v").build(); + assert!(matches!( + adapter.speech(req).await.unwrap_err(), + AibridgeError::UnsupportedCapability { .. } + )); + } + + // ============ resolve_speech_model ============ + + #[test] + fn resolve_speech_model_best_passthrough() { + assert_eq!(AssemblyAiAdapter::resolve_speech_model("best"), "best"); + } + + #[test] + fn resolve_speech_model_nano_passthrough() { + assert_eq!(AssemblyAiAdapter::resolve_speech_model("nano"), "nano"); + } + + #[test] + fn resolve_speech_model_unknown_defaults_best() { + assert_eq!(AssemblyAiAdapter::resolve_speech_model("whisper"), "best"); + assert_eq!(AssemblyAiAdapter::resolve_speech_model(""), "best"); + } + + // ============ build_transcript_payload ============ + + #[test] + fn build_payload_minimal_defaults() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])).build(); + let payload = + AssemblyAiAdapter::build_transcript_payload(&req, "https://up.example.com/u", "best"); + assert_eq!(payload["audio_url"], "https://up.example.com/u"); + assert_eq!(payload["speech_model"], "best"); + assert_eq!(payload["punctuate"], true); + assert_eq!(payload["format_text"], true); + // 无可选字段时不应出现 + assert!(payload.get("language_code").is_none()); + assert!(payload.get("speaker_labels").is_none()); + assert!(payload.get("redact_pii").is_none()); + assert!(payload.get("word_boost").is_none()); + } + + #[test] + fn build_payload_punctuate_format_text_override() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("punctuate", false) + .extra("format_text", false) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["punctuate"], false); + assert_eq!(payload["format_text"], false); + } + + #[test] + fn build_payload_with_language() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .language("zh") + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["language_code"], "zh"); + } + + #[test] + fn build_payload_language_from_extra_language_code() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("language_code", "ja") + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["language_code"], "ja"); + } + + #[test] + fn build_payload_req_language_takes_priority_over_extra() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .language("en") + .extra("language_code", "ja") + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["language_code"], "en"); + } + + #[test] + fn build_payload_with_speaker_labels() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("speaker_labels", true) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["speaker_labels"], true); + } + + #[test] + fn build_payload_with_all_bool_flags() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("speaker_labels", true) + .extra("filter_profanity", true) + .extra("sentiment_analysis", true) + .extra("auto_chapters", true) + .extra("entity_detection", true) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["speaker_labels"], true); + assert_eq!(payload["filter_profanity"], true); + assert_eq!(payload["sentiment_analysis"], true); + assert_eq!(payload["auto_chapters"], true); + assert_eq!(payload["entity_detection"], true); + } + + #[test] + fn build_payload_bool_flag_false_omitted() { + // 显式 false 不加入(与 Python `if kwargs.get(...)` 语义一致) + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("speaker_labels", false) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert!(payload.get("speaker_labels").is_none()); + } + + #[test] + fn build_payload_with_redact_pii() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("redact_pii", true) + .extra("redact_pii_policies", json!(["person_name", "email"])) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["redact_pii"], true); + assert_eq!( + payload["redact_pii_policies"], + json!(["person_name", "email"]) + ); + } + + #[test] + fn build_payload_redact_pii_default_empty_policies() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("redact_pii", true) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["redact_pii"], true); + assert_eq!(payload["redact_pii_policies"], json!([])); + } + + #[test] + fn build_payload_with_word_boost() { + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("word_boost", json!(["Apple", "Google"])) + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert_eq!(payload["word_boost"], json!(["Apple", "Google"])); + } + + #[test] + fn build_payload_word_boost_non_array_omitted() { + // 非 JSON 数组的 word_boost 不透传(与 Python isinstance 检查一致) + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("word_boost", "Apple") + .build(); + let payload = AssemblyAiAdapter::build_transcript_payload(&req, "https://u", "best"); + assert!(payload.get("word_boost").is_none()); + } + + // ============ parse_response ============ + + #[test] + fn parse_response_simple_text() { + let v = json!({ + "status": "completed", + "text": "hello world", + "language_code": "en", + "audio_duration": 5.5 + }); + let r = AssemblyAiAdapter::parse_response(&v, "best", "transcribe"); + assert_eq!(r.text, "hello world"); + assert_eq!(r.language.as_deref(), Some("en")); + assert!((r.duration.unwrap() - 5.5).abs() < f64::EPSILON); + assert!(r.segments.is_none()); + assert!(r.words.is_none()); + assert_eq!(r.task, "transcribe"); + assert_eq!(r.model.as_deref(), Some("best")); + } + + #[test] + fn parse_response_with_utterances_segments_and_words() { + let v = json!({ + "text": "hello world", + "utterances": [ + { + "start": 0, + "end": 1500, + "text": "hello", + "confidence": 0.95, + "speaker": "A", + "words": [ + {"text": "hello", "start": 0, "end": 500, "confidence": 0.99} + ] + }, + { + "start": 1500, + "end": 3000, + "text": "world", + "confidence": 0.9, + "speaker": "B" + } + ] + }); + let r = AssemblyAiAdapter::parse_response(&v, "best", "transcribe"); + assert_eq!(r.text, "hello world"); + let segs = r.segments.unwrap(); + assert_eq!(segs.len(), 2); + assert_eq!(segs[0].id, 0); + assert!((segs[0].start - 0.0).abs() < f64::EPSILON); + assert!((segs[0].end - 1.5).abs() < f64::EPSILON); + assert_eq!(segs[0].text, "hello"); + assert!((segs[0].confidence.unwrap() - 0.95).abs() < f64::EPSILON); + assert_eq!(segs[0].speaker.as_deref(), Some("A")); + assert_eq!(segs[1].speaker.as_deref(), Some("B")); + // 词级时间戳(来自 utterance[0].words) + let words = r.words.unwrap(); + assert_eq!(words.len(), 1); + assert_eq!(words[0].word, "hello"); + assert!((words[0].end - 0.5).abs() < f64::EPSILON); + } + + #[test] + fn parse_response_top_level_words_when_no_utterances() { + let v = json!({ + "text": "hi", + "words": [ + {"text": "hi", "start": 0, "end": 200, "confidence": 0.9} + ] + }); + let r = AssemblyAiAdapter::parse_response(&v, "nano", "transcribe"); + assert!(r.segments.is_none()); + let words = r.words.unwrap(); + assert_eq!(words.len(), 1); + assert!((words[0].start - 0.0).abs() < f64::EPSILON); + assert!((words[0].end - 0.2).abs() < f64::EPSILON); + assert_eq!(r.model.as_deref(), Some("nano")); + } + + #[test] + fn parse_response_translate_task() { + let v = json!({"text": "translated text"}); + let r = AssemblyAiAdapter::parse_response(&v, "best", "translate"); + assert_eq!(r.task, "translate"); + } + + #[test] + fn parse_response_missing_fields_defaults() { + let v = json!({}); + let r = AssemblyAiAdapter::parse_response(&v, "best", "transcribe"); + assert_eq!(r.text, ""); + assert!(r.language.is_none()); + assert!(r.duration.is_none()); + assert!(r.segments.is_none()); + assert!(r.words.is_none()); + } + + #[test] + fn parse_response_empty_utterances_falls_back_to_top_words() { + let v = json!({ + "text": "hi", + "utterances": [], + "words": [{"text": "hi", "start": 0, "end": 100}] + }); + let r = AssemblyAiAdapter::parse_response(&v, "best", "transcribe"); + // 空 utterances 数组当作无分段 + assert!(r.segments.is_none()); + // 回退到顶层 words(因为 utterances 为空数组,as_array 返 Some([]),不走 else 分支) + // 注意:空 utterances 时 segs 为空,words 也为空(来自 utterances 循环) + assert!(r.words.is_none()); + } + + // ============ ms_to_secs ============ + + #[test] + fn ms_to_secs_converts_milliseconds() { + assert!((ms_to_secs(Some(&json!(1000))) - 1.0).abs() < f64::EPSILON); + assert!((ms_to_secs(Some(&json!(2500))) - 2.5).abs() < f64::EPSILON); + assert!((ms_to_secs(Some(&json!(0))) - 0.0).abs() < f64::EPSILON); + } + + #[test] + fn ms_to_secs_none_returns_zero() { + assert!((ms_to_secs(None) - 0.0).abs() < f64::EPSILON); + } + + // ============ 错误映射 ============ + + #[test] + fn parse_error_message_from_error_field() { + let body = r#"{"error":"invalid api key"}"#; + assert_eq!(parse_assemblyai_error_message(body, 401), "invalid api key"); + } + + #[test] + fn parse_error_message_from_message_field() { + let body = r#"{"message":"something went wrong"}"#; + assert_eq!( + parse_assemblyai_error_message(body, 500), + "something went wrong" + ); + } + + #[test] + fn parse_error_message_fallback_to_body_text() { + assert_eq!( + parse_assemblyai_error_message("plain text error", 500), + "plain text error" + ); + } + + #[test] + fn parse_error_message_empty_body_falls_back_to_http_status() { + assert_eq!(parse_assemblyai_error_message("", 500), "HTTP 500"); + assert_eq!(parse_assemblyai_error_message(" ", 500), "HTTP 500"); + } + + #[test] + fn map_error_401_is_authentication() { + let err = map_assemblyai_error(401, r#"{"error":"unauthorized"}"#); + assert!(matches!(err, AibridgeError::Authentication { .. })); + assert!(!err.is_retryable()); + } + + #[test] + fn map_error_403_is_authentication() { + let err = map_assemblyai_error(403, "forbidden"); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[test] + fn map_error_429_is_rate_limit() { + let err = map_assemblyai_error(429, r#"{"error":"too many requests"}"#); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + assert!(err.is_retryable()); + } + + #[test] + fn map_error_400_is_validation() { + let err = map_assemblyai_error(400, r#"{"error":"bad audio_url"}"#); + assert!(matches!(err, AibridgeError::Validation { .. })); + assert!(!err.is_retryable()); + } + + #[test] + fn map_error_500_is_api_and_retryable() { + let err = map_assemblyai_error(500, r#"{"error":"internal"}"#); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + assert!(err.is_retryable()); + } + + #[test] + fn map_error_other_4xx_is_api() { + let err = map_assemblyai_error(418, "teapot"); + assert!(matches!(err, AibridgeError::Api { status: 418, .. })); + assert!(!err.is_retryable()); + } + + // ============ transcribe 完整流程(HTTP,mockito) ============ + + #[tokio::test] + async fn transcribe_full_flow_bytes_input() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .match_header("Authorization", "test-key") + .match_header("Content-Type", "application/octet-stream") + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u1"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .match_header("Authorization", "test-key") + .with_status(200) + .with_body(r#"{"id":"t-123"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-123") + .match_header("Authorization", "test-key") + .with_status(200) + .with_body( + r#"{"status":"completed","text":"hello world","language_code":"en","audio_duration":5.5}"#, + ) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1, 2, 3])) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "hello world"); + assert_eq!(result.language.as_deref(), Some("en")); + assert!((result.duration.unwrap() - 5.5).abs() < f64::EPSILON); + assert_eq!(result.task, "transcribe"); + assert_eq!(result.model.as_deref(), Some("best")); + } + + #[tokio::test] + async fn transcribe_nano_model_resolved_in_result() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .match_body(mockito::Matcher::Json(json!({ + "audio_url": "https://up.example.com/u", + "speech_model": "nano", + "punctuate": true, + "format_text": true + }))) + .with_status(200) + .with_body(r#"{"id":"t-nano"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-nano") + .with_status(200) + .with_body(r#"{"status":"completed","text":"nano result"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("nano", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "nano result"); + assert_eq!(result.model.as_deref(), Some("nano")); + } + + #[tokio::test] + async fn transcribe_unknown_model_defaults_best() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + let create_mock = server + .mock("POST", TRANSCRIPT_PATH) + .match_body(mockito::Matcher::PartialJson( + json!({"speech_model": "best"}), + )) + .with_status(200) + .with_body(r#"{"id":"t-b"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-b") + .with_status(200) + .with_body(r#"{"status":"completed","text":"ok"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("unknown-model", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.model.as_deref(), Some("best")); + create_mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_url_input_skips_upload() { + let mut server = mockito::Server::new_async().await; + // upload 不应被调用 + let upload_mock = server + .mock("POST", UPLOAD_PATH) + .expect(0) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .match_body(mockito::Matcher::PartialJson( + json!({"audio_url": "https://example.com/audio.mp3"}), + )) + .with_status(200) + .with_body(r#"{"id":"t-url"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-url") + .with_status(200) + .with_body(r#"{"status":"completed","text":"url audio"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = + TranscribeRequest::builder("best", FileInput::url("https://example.com/audio.mp3")) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "url audio"); + upload_mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_extra_audio_url_skips_upload() { + let mut server = mockito::Server::new_async().await; + let upload_mock = server + .mock("POST", UPLOAD_PATH) + .expect(0) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-x"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-x") + .with_status(200) + .with_body(r#"{"status":"completed","text":"ok"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + // 即使 file 是 Bytes,extra.audio_url 优先,跳过上传 + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("audio_url", "https://cdn.example.com/a.mp3") + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "ok"); + upload_mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_with_speaker_labels() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .match_body(mockito::Matcher::PartialJson( + json!({"speaker_labels": true}), + )) + .with_status(200) + .with_body(r#"{"id":"t-sl"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-sl") + .with_status(200) + .with_body( + r#"{"status":"completed","text":"hi","utterances":[{"start":0,"end":1000,"text":"hi","confidence":0.9,"speaker":"A"}]}"#, + ) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("speaker_labels", true) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "hi"); + let segs = result.segments.unwrap(); + assert_eq!(segs.len(), 1); + assert_eq!(segs[0].speaker.as_deref(), Some("A")); + assert!((segs[0].end - 1.0).abs() < f64::EPSILON); + } + + #[tokio::test] + async fn transcribe_processing_then_completed() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-p1"}"#) + .create_async() + .await; + // 第一次轮询返 processing(FIFO:先创建先匹配) + let processing_mock = server + .mock("GET", "/transcript/t-p1") + .with_status(200) + .with_body(r#"{"status":"processing"}"#) + .expect(1) + .create_async() + .await; + // 第二次轮询返 completed + let completed_mock = server + .mock("GET", "/transcript/t-p1") + .with_status(200) + .with_body(r#"{"status":"completed","text":"after poll"}"#) + .expect(1) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .extra("max_polls", 5u64) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "after poll"); + processing_mock.assert_async().await; + completed_mock.assert_async().await; + } + + #[tokio::test] + async fn transcribe_error_status_returns_api_error() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-e1"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-e1") + .with_status(200) + .with_body(r#"{"status":"error","error":"audio too short"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + match err { + AibridgeError::Api { message, .. } => { + assert!(message.contains("audio too short")); + } + other => panic!("应为 Api 错误,实际: {other:?}"), + } + } + + #[tokio::test] + async fn transcribe_translate_sets_task() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-tr"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-tr") + .with_status(200) + .with_body(r#"{"status":"completed","text":"translated","language_code":"en"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .translate(true) + .extra("polling_interval", 0.001) + .build(); + let result = adapter.transcribe(req).await.unwrap(); + assert_eq!(result.text, "translated"); + assert_eq!(result.task, "translate"); + } + + // ============ transcribe 错误路径 ============ + + #[tokio::test] + async fn transcribe_upload_401_returns_authentication() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(401) + .with_body(r#"{"error":"invalid api key"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Authentication { .. })); + } + + #[tokio::test] + async fn transcribe_upload_429_returns_rate_limit() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(429) + .with_body(r#"{"error":"too many requests"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::RateLimit { .. })); + } + + #[tokio::test] + async fn transcribe_create_400_returns_validation() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(400) + .with_body(r#"{"error":"invalid audio_url"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Validation { .. })); + } + + #[tokio::test] + async fn transcribe_create_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(500) + .with_body(r#"{"error":"internal"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[tokio::test] + async fn transcribe_poll_500_returns_api() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-s5"}"#) + .create_async() + .await; + server + .mock("GET", "/transcript/t-s5") + .with_status(500) + .with_body(r#"{"error":"server error"}"#) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Api { status: 500, .. })); + } + + #[tokio::test] + async fn transcribe_max_polls_exhausted_returns_timeout() { + let mut server = mockito::Server::new_async().await; + server + .mock("POST", UPLOAD_PATH) + .with_status(200) + .with_body(r#"{"upload_url":"https://up.example.com/u"}"#) + .create_async() + .await; + server + .mock("POST", TRANSCRIPT_PATH) + .with_status(200) + .with_body(r#"{"id":"t-to"}"#) + .create_async() + .await; + // 始终返 processing,模拟永不完成 + server + .mock("GET", "/transcript/t-to") + .with_status(200) + .with_body(r#"{"status":"processing"}"#) + .expect(2) + .create_async() + .await; + + let adapter = make_adapter(Some(server.url())); + let req = TranscribeRequest::builder("best", FileInput::bytes(vec![1])) + .extra("polling_interval", 0.001) + .extra("max_polls", 2u64) + .build(); + let err = adapter.transcribe(req).await.unwrap_err(); + assert!(matches!(err, AibridgeError::Timeout)); + } +} From 12925adcb51b753fe0adc434a21bd16b1120cd87 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 02:09:08 +0800 Subject: [PATCH 47/55] =?UTF-8?q?feat(aibridge-core):=20=E9=98=B6=E6=AE=B5?= =?UTF-8?q?2c=20=E7=AC=AC=E4=B8=89=E6=89=B9=E6=94=B6=E5=B0=BE=20=E6=B3=A8?= =?UTF-8?q?=E5=86=8C=20deepgram=20+=20assemblyai=20=E5=88=B0=E5=B7=A5?= =?UTF-8?q?=E5=8E=82=EF=BC=88=E9=98=B6=E6=AE=B52=20=E5=85=A8=E9=83=A8?= =?UTF-8?q?=E5=AE=8C=E6=88=90=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- crates/aibridge-core/src/adapter/factory.rs | 52 +++++++++++++++------ crates/aibridge-core/src/adapters/mod.rs | 6 +++ 2 files changed, 45 insertions(+), 13 deletions(-) diff --git a/crates/aibridge-core/src/adapter/factory.rs b/crates/aibridge-core/src/adapter/factory.rs index 94fb03f..fc98c80 100644 --- a/crates/aibridge-core/src/adapter/factory.rs +++ b/crates/aibridge-core/src/adapter/factory.rs @@ -16,11 +16,13 @@ use crate::adapters::aggregation_platforms::{ }; use crate::adapters::agnes::AgnesAdapter; use crate::adapters::anthropic::AnthropicAdapter; +use crate::adapters::assemblyai::AssemblyAiAdapter; use crate::adapters::azure::AzureAdapter; use crate::adapters::cartesia::CartesiaAdapter; use crate::adapters::chinese::{ DoubaoAdapter, ErnieAdapter, KimiAdapter, MiniMaxAdapter, QwenAdapter, ZhipuAdapter, }; +use crate::adapters::deepgram::DeepgramAdapter; use crate::adapters::echo::EchoAdapter; use crate::adapters::edge_tts::EdgeTtsAdapter; use crate::adapters::elevenlabs::ElevenLabsAdapter; @@ -87,7 +89,6 @@ pub const KNOWN_PROVIDERS: &[&str] = &[ "edge-tts", "elevenlabs", "cartesia", - // 阶段 2c 待实现: "deepgram", "assemblyai", ]; @@ -155,10 +156,11 @@ pub fn create_adapter(config: ProviderConfig) -> Result> { "elevenlabs" | "eleven" | "11labs" => Ok(Box::new(ElevenLabsAdapter::new(config)?)), // (cartesia / sonic 均指向 CartesiaAdapter) "cartesia" | "sonic" => Ok(Box::new(CartesiaAdapter::new(config)?)), - // 阶段 2c 适配器占位 - "deepgram" | "assemblyai" => Err(AibridgeError::ProviderNotFound { - provider: format!("{provider}(阶段 2c 待实现)"), - }), + // 阶段 2c 音频 ASR:别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用 + // (deepgram / dg 均指向 DeepgramAdapter) + "deepgram" | "dg" => Ok(Box::new(DeepgramAdapter::new(config)?)), + // (assemblyai / assembly / aai 均指向 AssemblyAiAdapter) + "assemblyai" | "assembly" | "aai" => Ok(Box::new(AssemblyAiAdapter::new(config)?)), // 未知 provider _ => Err(AibridgeError::provider_not_found(format!( "{provider}(未知 provider,支持:{})", @@ -541,14 +543,37 @@ mod tests { } #[test] - fn create_phase2_adapter_returns_phase2_message() { - // deepgram 仍为阶段 2c 占位(未实现),返 ProviderNotFound - let result = create_adapter(config_for("deepgram")); - if let Err(AibridgeError::ProviderNotFound { provider }) = result { - assert!(provider.contains("阶段 2")); - } else { - panic!("应为 ProviderNotFound"); - } + fn create_deepgram_returns_adapter() { + // 阶段 2c:DeepgramAdapter 自带 DEFAULT_API_BASE 兜底,仅需 api_key + let adapter = create_adapter(config_for("deepgram")).expect("工厂应能创建 deepgram 适配器"); + assert_eq!(adapter.provider_type(), "deepgram"); + } + + #[test] + fn create_assemblyai_returns_adapter() { + // 阶段 2c:AssemblyAiAdapter 自带 DEFAULT_API_BASE 兜底,api_key 可空(调用时 401) + let adapter = + create_adapter(config_for("assemblyai")).expect("工厂应能创建 assemblyai 适配器"); + assert_eq!(adapter.provider_type(), "assemblyai"); + } + + #[test] + fn create_deepgram_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用: + // dg -> deepgram(指向 DeepgramAdapter) + let dg = create_adapter(config_for("dg")).expect("别名 dg 应映射到 deepgram"); + assert_eq!(dg.provider_type(), "deepgram"); + } + + #[test] + fn create_assemblyai_aliases_map_to_main_provider_type() { + // 别名对齐 Python agn/adapters/audio_adapters.py 末尾 register 调用: + // assembly / aai -> assemblyai(指向 AssemblyAiAdapter) + let assembly = + create_adapter(config_for("assembly")).expect("别名 assembly 应映射到 assemblyai"); + assert_eq!(assembly.provider_type(), "assemblyai"); + let aai = create_adapter(config_for("aai")).expect("别名 aai 应映射到 assemblyai"); + assert_eq!(aai.provider_type(), "assemblyai"); } #[test] @@ -558,6 +583,7 @@ mod tests { assert!(is_known_provider("edge-tts")); assert!(is_known_provider("elevenlabs")); assert!(is_known_provider("cartesia")); + assert!(is_known_provider("deepgram")); assert!(is_known_provider("assemblyai")); assert!(is_known_provider("kling")); // 阶段 2a 已实现 provider 应被识别 diff --git a/crates/aibridge-core/src/adapters/mod.rs b/crates/aibridge-core/src/adapters/mod.rs index 7851c84..875c75e 100644 --- a/crates/aibridge-core/src/adapters/mod.rs +++ b/crates/aibridge-core/src/adapters/mod.rs @@ -69,3 +69,9 @@ pub mod elevenlabs; /// Cartesia 适配器:阶段 2c 音频,Cartesia Sonic TTS 文字转语音(低延迟流式) pub mod cartesia; + +/// Deepgram 适配器:阶段 2c 音频,Deepgram ASR 语音转文字(Token 鉴权 / REST 协议) +pub mod deepgram; + +/// AssemblyAI 适配器:阶段 2c 音频,AssemblyAI ASR 语音转文字(Key 鉴权 / REST 协议) +pub mod assemblyai; From 141ec2e409cb0fa4b905f401e50b2a42f7d95c50 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 04:08:55 +0800 Subject: [PATCH 48/55] =?UTF-8?q?docs:=20=E9=98=B6=E6=AE=B53=20Python=20v1?= =?UTF-8?q?=E2=86=92v2=20=E8=BF=81=E7=A7=BB=E6=8C=87=E5=8D=97=20+=20README?= =?UTF-8?q?=20+=20=E8=BF=9B=E5=BA=A6=E6=9B=B4=E6=96=B0?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit --- README_aibridge.md | 410 +++++++++++++++++++++ docs/PROGRESS.md | 146 +++++--- docs/migration-guide.md | 766 ++++++++++++++++++++++++++++++++++++++++ 3 files changed, 1280 insertions(+), 42 deletions(-) create mode 100644 README_aibridge.md create mode 100644 docs/migration-guide.md diff --git a/README_aibridge.md b/README_aibridge.md new file mode 100644 index 0000000..bdf8cd4 --- /dev/null +++ b/README_aibridge.md @@ -0,0 +1,410 @@ +# AIBridge + +> 跨语言 AI 统一接口 SDK —— 一套 API 调用所有 AI 模型。 +> Rust 核心 + Python / JS-TS / Go / JVM / .NET 五语言原生绑定。 + +[![状态](https://img.shields.io/badge/状态-阶段3发布收尾中-yellow)](docs/PROGRESS.md) +[![provider](https://img.shields.io/badge/provider-38+-blue)](#支持的-provider) +[![核心测试](https://img.shields.io/badge/core%20单测-1448-brightgreen)](docs/PROGRESS.md) + +AIBridge(原名 agn-sdk)是多模态 AI 统一接口 SDK:文本对话(chat)、图像生成(image)、视频生成(video)、文字转语音(TTS)、语音转文字(ASR)、文本嵌入(embed),一套 API 调用 38+ 个 AI provider。 + +- **五语言原生 import**:Python / JS-TS / Go / JVM(Java,Kotlin) / .NET(C#) 直接调用同一套能力 +- **Rust 核心**:无 GIL、原生 async、零成本抽象,IO 密集场景高性能 +- **38+ provider**:OpenAI / Claude / Gemini / 通义千问 / 智谱 / 文心一言 / DeepSeek / 火山引擎 / Stability / Runway / Pika / 可灵 / Edge-TTS / ElevenLabs / Deepgram … +- **六大能力**:chat(含流式)/ image / video / TTS / ASR / embed +- **免认证 provider**:edge-tts 免费 TTS,无需 API key + +> **v2 是破坏性升级**。若你已在用 Python v1(`agn-sdk`),请参阅 [迁移指南](docs/migration-guide.md)。 + +--- + +## 目录 + +- [架构](#架构) +- [支持的 provider](#支持的-provider) +- [五语言快速开始](#五语言快速开始) + - [Python](#python) + - [Node.js / TypeScript](#nodejs--typescript) + - [Go](#go) + - [Java / Kotlin (JVM)](#java--kotlin-jvm) + - [C# / .NET](#c--net) +- [安装](#安装) +- [能力一览](#能力一览) +- [相关文档](#相关文档) + +--- + +## 架构 + +``` + aibridge-core (Rust, 纯 async 逻辑) + ┌──────────┴──────────┐ + 直连(原生async) C ABI (aibridge-ffi cdylib) + ┌─────┴─────┐ ┌─────┬─────┬─────┐ + aibridge- aibridge- aibridge- aibridge- aibridge- + python node go jvm dotnet + (PyO3) (napi-rs) (CGO) (JNA) (P/Invoke) + asyncio Promise/ goroutine CompletableFuture Task/ + AsyncIter AsyncIter +channel /Flow IAsyncEnum +``` + +- **Python / JS-TS 直连 Rust 核心**:PyO3 / napi-rs 直连 `aibridge-core`,无 JSON 序列化边界,享真正原生 async。 +- **Go / JVM / .NET 走 C ABI**:通过 `aibridge-ffi` 的 C ABI(句柄 + JSON 边界 + 全局 tokio runtime),各语言用原生异步原语包装。 +- **五种语言共享同一个 Rust 核心**,绑定层都薄,行为一致。 + +详见 [设计文档](docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md)。 + +--- + +## 支持的 provider + +共 **38 个真实 provider + 1 个 mock**(echo,测试用),按类别: + +| 类别 | provider | +|---|---| +| **MVP** | `openai` `agnes` `volcengine_cv`(火山引擎) `gemini` | +| **OpenAI 兼容族** | `azure` `siliconflow`(别名 `sf`) `togetherai`(`together`) `fireworksai`(`fireworks`) `cloudflareai`(`cloudflare` / `workersai`) | +| **扩展模型** | `grok`(`xaigrok`) `yi`(`lingyiwanwu`) `sensenova`(`shangtang`) `hunyuan`(`tencent_hunyuan`) `groq` | +| **更多模型** | `deepseek` `stepfun`(`step`) `mistral` `cohere` `perplexity` | +| **新兴模型** | `ideogram`(`ideo`) `luma`(`dream-machine` / `lumalabs`) `llama`(`meta-llama` / `meta`) | +| **中文模型** | `qwen`(通义千问) `zhipu`(智谱) `doubao`(豆包) `ernie`(文心一言) `kimi` `minimax` | +| **独立协议** | `anthropic`(Claude) `stability`(Stability AI) `runway` `pika` `kling`(可灵) | +| **音频 TTS** | `edge-tts`(免费,免认证,别名 `edge_tts` / `edge`) `elevenlabs`(`eleven` / `11labs`) `cartesia`(`sonic`) | +| **音频 ASR** | `deepgram`(`dg`) `assemblyai`(`assembly` / `aai`) | +| **Mock** | `echo`(免认证,测试用,不调网络) | + +--- + +## 五语言快速开始 + +下面每个示例都用 **`echo` 适配器**(免认证、不调网络),可直接运行验证管线。真实使用时把 `provider="echo"` 换成 `provider="openai"` 等并传入 `api_key`。 + +### Python + +```python +import asyncio +from aibridge import Client + +async def main(): + # echo 免认证;真实用例:Client(provider="openai", api_key="sk-xxx") + client = Client(provider="echo") + await client.start() + + # 文本对话 + resp = await client.chat( + model="echo-chat", + messages=[{"role": "user", "content": "你好"}], + ) + print(resp.choices[0].message.content) # "你好 [echo]" + + # 流式对话 + stream = await client.chat_stream( + model="echo-chat", + messages=[{"role": "user", "content": "你好"}], + ) + async for chunk in stream: + print(chunk.choices[0].content, end="", flush=True) + print() + + # 文字转语音 + audio = await client.speech(model="echo-tts", input="你好", voice="alloy") + print(f"音频 {len(audio.audio_data)} 字节") + + await client.close() + +asyncio.run(main()) +``` + +运行: + +```bash +pip install maturin +maturin develop -m crates/aibridge-python/Cargo.toml +python examples/hello_python.py +``` + +### Node.js / TypeScript + +```javascript +const { Client } = require('./crates/aibridge-node'); + +async function main() { + // echo 免认证;真实用例:new Client('openai', { apiKey: 'sk-xxx' }) + const client = new Client('echo', {}); + await client.start(); + + // 文本对话 + const resp = await client.chat({ + model: 'echo-chat', + messages: [{ role: 'user', content: '你好' }], + }); + console.log(resp.choices[0].message.content); // "你好 [echo]" + + // 流式对话 + const stream = await client.chatStream({ + model: 'echo-chat', + messages: [{ role: 'user', content: '你好' }], + }); + for await (const chunk of stream) { + const delta = chunk.choices[0]; + if (delta.content) process.stdout.write(delta.content); + } + console.log(); + + // 文字转语音 + const audio = await client.speech({ + model: 'echo-tts', + input: '你好', + voice: 'alloy', + }); + console.log(`音频 ${audio.audioData.length} 字节`); + + await client.close(); +} + +main(); +``` + +运行: + +```bash +cd crates/aibridge-node && npm install && napi build && cd ../.. +node examples/hello_node.js +``` + +### Go + +```go +package main + +import ( + "fmt" + + aibridge "github.com/aibridge/aibridge-go" +) + +func main() { + // echo 免认证;真实用例:aibridge.NewClient("openai", &aibridge.ClientOpts{ApiKey: "sk-xxx"}) + client, err := aibridge.NewClient("echo", nil) + if err != nil { + panic(err) + } + defer client.Close() + client.Start() + + // 文本对话 + chatReq := &aibridge.ChatRequest{ + Model: "echo-chat", + Messages: []aibridge.ChatMessage{aibridge.NewUserTextMessage("你好")}, + } + resp, err := client.Chat(chatReq) + if err != nil { + panic(err) + } + fmt.Println(resp.Choices[0].Message.Content) // "你好 [echo]" + + // 流式对话 + stream, err := client.ChatStream(chatReq) + if err != nil { + panic(err) + } + for chunk := range stream.Ch() { + if len(chunk.Choices) > 0 { + fmt.Print(chunk.Choices[0].Delta.Content) + } + } + fmt.Println() + + // 文字转语音 + speechReq := &aibridge.SpeechRequest{ + Model: "echo-tts", + Input: "你好", + Voice: aibridge.SingleVoice("alloy"), + } + speech, err := client.Speech(speechReq) + if err != nil { + panic(err) + } + fmt.Printf("音频 %d 字节\n", len(speech.AudioData)) +} +``` + +运行: + +```bash +cargo build -p aibridge-ffi +cd bindings/go +CGO_ENABLED=1 DYLD_LIBRARY_PATH=../../target/debug go run ./example +``` + +### Java / Kotlin (JVM) + +```java +package io.aibridge; + +import java.util.List; + +public class Hello { + public static void main(String[] args) { + // echo 免认证;真实用例:new Client("openai", "sk-xxx") + try (Client client = new Client("echo")) { + client.start(); + + // 文本对话 + ChatRequest req = ChatRequest.builder( + "echo-chat", + List.of(ChatMessage.user("你好"))) + .build(); + ChatCompletion resp = client.chat(req); + System.out.println(resp.choices.get(0).message.content); // "你好 [echo]" + + // 流式对话 + ChatRequest streamReq = ChatRequest.builder( + "echo-chat", + List.of(ChatMessage.user("你好"))) + .stream(true) + .build(); + try (ChatStream stream = client.chatStream(streamReq)) { + while (stream.hasNext()) { + ChatCompletionChunk chunk = stream.next(); + String c = chunk.firstDeltaContent(); + if (c != null) System.out.print(c); + } + } + System.out.println(); + + // 文字转语音 + SpeechRequest speechReq = SpeechRequest.builder("echo-tts", "你好", "alloy").build(); + SpeechResultFull audio = client.speech(speechReq); + System.out.println("音频 " + audio.audioLength() + " 字节"); + } + } +} +``` + +运行: + +```bash +cd bindings/jvm && ./gradlew run +``` + +### C# / .NET + +```csharp +using AIBridge; + +// echo 免认证;真实用例:new Client("openai", "sk-xxx") +using var client = new Client("echo"); +client.Start(); + +// 文本对话 +var chatReq = new ChatRequest("echo-chat", new[] +{ + ChatMessage.User("你好"), +}); +ChatCompletion resp = client.Chat(chatReq); +Console.WriteLine(resp.Choices[0].Message.Content); // "你好 [echo]" + +// 流式对话 +int chunkCount = 0; +var assembled = new System.Text.StringBuilder(); +await foreach (ChatCompletionChunk chunk in client.ChatStreamAsync(chatReq)) +{ + chunkCount++; + if (chunk.Choices.Count > 0 && chunk.Choices[0].Delta.Content != null) + assembled.Append(chunk.Choices[0].Delta.Content); +} +Console.WriteLine(assembled.ToString()); + +// 文字转语音 +var speechReq = new SpeechRequest("echo-tts", "你好", "alloy"); +SpeechResult audio = client.Speech(speechReq); +Console.WriteLine($"音频 {audio.AudioData.Length} 字节"); +``` + +运行: + +```bash +cargo build -p aibridge-ffi +cd bindings/dotnet && dotnet run +``` + +--- + +## 安装 + +> 阶段 3 发布进行中,以下为各语言的目标安装方式。发布前可从源码构建。 + +| 语言 | 安装命令 | 包名 | +|---|---|---| +| Python | `pip install aibridge` | PyPI `aibridge` | +| Node.js | `npm install aibridge` | npm `aibridge` | +| Go | `go get github.com/aibridge/aibridge-go`(需单独装 libaibridge) | Go module `aibridge-go` | +| JVM | Maven `io.aibridge:aibridge` | Maven Central | +| .NET | `dotnet add package AIBridge` | NuGet `AIBridge` | + +从源码构建(开发/发布前): + +```bash +# Rust 核心 + ffi +cargo build --workspace + +# Python 绑定 +pip install maturin +maturin develop -m crates/aibridge-python/Cargo.toml + +# Node 绑定 +cd crates/aibridge-node && npm install && napi build + +# Go / JVM / .NET 绑定需先 cargo build -p aibridge-ffi 产 libaibridge 动态库 +``` + +--- + +## 能力一览 + +| 能力 | 方法 | 说明 | +|---|---|---| +| 文本对话 | `chat` | 支持多轮、system/user/assistant/tool、多模态、工具调用 | +| 流式对话 | `chat_stream` | 原生异步迭代器,逐块产出 | +| 图像生成 | `image_generate` | 文生图、图生图(reference_images)、局部重绘(mask) | +| 视频生成 | `video_create` + `video_poll` | 文生视频、图生视频,任务轮询 | +| 文字转语音 | `speech` | TTS,支持音色候选列表自动降级 | +| 语音转文字 | `transcribe` | ASR,支持文件路径/URL/bytes/base64 输入 | +| 文本嵌入 | `embed` | 文本向量化 | +| 模型列表 | `list_models` | 实时拉取 provider 可用模型 | +| 音色列表 | `list_voices` / `recommend_voices` | 音色健康检查/推荐/自动降级 | + +### 错误处理 + +统一错误基类 `AibridgeError`,子类按错误性质分类(各语言异常名一致): + +- `AuthenticationError` 认证失败 +- `RateLimitError` 限流(含 `retry_after`) +- `ValidationError` 参数校验 +- `ModelNotFoundError` 模型不存在 +- `APIError` Provider API 错误 +- `NetworkError` 网络错误 +- `TimeoutError` 超时 +- `UnsupportedCapabilityError` 能力不支持 +- `ProviderNotFoundError` provider 不存在 +- `VoiceNotAvailableError` 音色不可用 +- `ServiceUnavailableError` 服务暂不可用(可重试) + +错误带稳定 `code`(snake_case,如 `rate_limit_error`)与 `retryable` 标识,便于业务层重试决策。 + +--- + +## 相关文档 + +- [设计文档](docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md) — 架构、数据模型、FFI 边界、异步桥接、错误处理、适配器迁移策略 +- [迁移指南](docs/migration-guide.md) — Python v1(agn-sdk)→ v2(aibridge)破坏性升级对照与示例 +- [进度文档](docs/PROGRESS.md) — 当前实施进度与接手指南 +- [原 README(v1)](README.md) — Python v1 文档(归档参考) + +--- + +## License + +同仓库主 LICENSE。 diff --git a/docs/PROGRESS.md b/docs/PROGRESS.md index 7b1b30d..0c49b22 100644 --- a/docs/PROGRESS.md +++ b/docs/PROGRESS.md @@ -1,7 +1,7 @@ # AIBridge Rust 重构 · 进度与接手文档 > 本文档供任何 agent 接手 AIBridge Rust 重构工作使用。自包含,不依赖 Claude memory。 -> 最后更新:2026-07-07 +> 最后更新:2026-07-08 --- @@ -19,6 +19,8 @@ |---|---|---| | 设计文档 | [docs/superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md](superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md) | 架构、数据模型、FFI 边界、异步桥接、错误处理、适配器迁移策略 | | 实现计划 | [docs/superpowers/plans/2026-07-07-aibridge-implementation-plan.md](superpowers/plans/2026-07-07-aibridge-implementation-plan.md) | 阶段 0-3 任务分解、多 agent 编排策略、里程碑 | +| 迁移指南 | [docs/migration-guide.md](migration-guide.md) | Python v1(agn-sdk)→ v2(aibridge)破坏性升级对照与示例 | +| v2 README | [README_aibridge.md](../README_aibridge.md) | 五语言快速开始 + provider 列表 | | 本进度文档 | docs/PROGRESS.md | 当前进度 + 接手指南(本文档) | ## 3. 架构概览 @@ -45,11 +47,11 @@ aibridge/ │ ├── go/ # CGO 调 ffi │ ├── jvm/ # JNA 调 ffi(Java) │ └── dotnet/ # P/Invoke 调 ffi(C#) -├── docs/ # 设计文档 + 计划 + 本进度文档 +├── docs/ # 设计文档 + 计划 + 迁移指南 + 本进度文档 └── examples/ # 五语言 hello world(echo adapter) ``` -## 4. 当前进度(截至 2026-07-07) +## 4. 当前进度(截至 2026-07-08) ### ✅ 阶段 0:地基(完成) - Cargo workspace + 4 crate 骨架 @@ -67,7 +69,11 @@ aibridge/ - 错误 code 统一(对齐 Rust error.rs 的 code() 值) - dylib 产物名冲突修复(aibridge-python lib 改 `_aibridge`) -### ✅ 阶段 2a:OpenAI 兼容族 + 部分独立协议(完成,23 provider) +### ✅ 阶段 2:全量适配器迁移(完成,38 真实 provider + echo mock) + +阶段 2 分三批搬完 v1 全部适配器,编译期 match 工厂已注册全部 provider。工厂占位分支已清空。 + +#### 阶段 2a:OpenAI 兼容族(完成,24 provider) | 文件 | provider(含别名) | |---|---| | openai.rs | openai | @@ -81,22 +87,47 @@ aibridge/ | emerging_models.rs | ideogram\|ideo, luma\|dream-machine\|lumalabs, llama\|meta-llama\|meta | | chinese.rs | qwen, zhipu, doubao, ernie, kimi, minimax | -### ⏳ 阶段 2b:独立协议 5 个(待做) -anthropic / stability / runway / pika / kling - -### ⏳ 阶段 2c:音频 5 个(待做,二进制载荷) -edge-tts / elevenlabs / cartesia / deepgram / assemblyai - -### ⏳ 阶段 3:发布(待做) -- CI 矩阵(平台 × 语言,交叉编译) -- 五语言包发布(PyPI `aibridge` / npm `aibridge` / Maven `io.aibridge:aibridge` / NuGet `AIBridge` / Go module `aibridge-go`) -- Python v1→v2 迁移指南 -- 文档网站 -- 旧版 v1 归档 + 打 v2.0.0 tag +#### 阶段 2b:独立协议 5 个(完成) +| 文件 | provider | 能力 | +|---|---|---| +| anthropic.rs | anthropic | Claude messages API,流式 SSE,文本对话/多模态 | +| stability.rs | stability | Stability AI 文生图/图生图 | +| runway.rs | runway | 视频生成(文生视频/图生视频/任务轮询) | +| pika.rs | pika | 视频生成(文生视频/图生视频/任务轮询) | +| kling.rs | kling | 可灵视频生成(文生视频/图生视频/任务轮询) | + +#### 阶段 2c:音频 5 个(完成,二进制载荷) +| 文件 | provider | 能力 | 备注 | +|---|---|---|---| +| edge_tts.rs | edge-tts(别名 edge_tts / edge) | 免费 TTS | 免认证(`requires_api_key=false`) | +| elevenlabs.rs | elevenlabs(别名 eleven / 11labs) | TTS | 高质量音色/多语种/克隆 | +| cartesia.rs | cartesia(别名 sonic) | TTS | Sonic 低延迟流式 | +| deepgram.rs | deepgram(别名 dg) | ASR | Token 鉴权 / REST | +| assemblyai.rs | assemblyai(别名 assembly / aai) | ASR | Key 鉴权 / REST | + +#### 已迁移 provider 完整列表(38 真实 + 1 mock) + +**MVP(4)**:openai、agnes、volcengine_cv、gemini +**兼容族(24)**:azure、siliconflow、togetherai、fireworksai、cloudflareai、grok、yi、sensenova、hunyuan、groq、deepseek、stepfun、mistral、cohere、perplexity、ideogram、luma、llama、qwen、zhipu、doubao、ernie、kimi、minimax +**独立协议(5)**:anthropic、stability、runway、pika、kling +**音频(5)**:edge-tts、elevenlabs、cartesia、deepgram、assemblyai +**Mock(1)**:echo(阶段 0.6 管线验证用,常驻) + +### ⏳ 阶段 3:发布收尾(进行中) +- [x] Python v1→v2 迁移指南([docs/migration-guide.md](migration-guide.md)) +- [x] v2 README([README_aibridge.md](../README_aibridge.md)) +- [x] 进度文档更新(本文档) +- [ ] CI 矩阵(平台 × 语言,交叉编译) +- [ ] 五语言包发布(PyPI `aibridge` / npm `aibridge` / Maven `io.aibridge:aibridge` / NuGet `AIBridge` / Go module `aibridge-go`) +- [ ] 文档网站 +- [ ] 旧版 v1 归档 + 打 v2.0.0 tag +- [ ] .NET hello world 验证(待 dotnet sdk) +- [ ] 一致性测试纳入 CI(当前手动跑) +- [ ] 真实 API key 冒烟测试(用户验证) ## 5. 测试状态 -- **aibridge-core**:810 单测全通过 +- **aibridge-core**:1448 单测全通过(含 38 provider mock HTTP 测试 + 数据模型 + 错误 + 路由) - **aibridge-ffi**:39 单测全通过 - **五语言 hello world**(echo adapter):Python/Node/Go/JVM 跑通,.NET 代码就绪待 dotnet - **跨语言一致性测试**:tests/consistency/(四语言 chat/stream/speech/错误全一致) @@ -107,6 +138,7 @@ edge-tts / elevenlabs / cartesia / deepgram / assemblyai 2. **一致性测试纳入 CI**:当前手动跑,待接入 CI matrix 3. **真实 IO 流式验证**:Python/Node 流式重构已完成(不阻塞推理),但真实 API key 验证待用户 4. **dylib 分发**:Go/JVM/.NET 依赖 libaibridge,JVM/.NET 打进包,Go 提供安装脚本(阶段 3 处理) +5. **Python 绑定能力暴露**:Rust 核心已全部实现 38 provider + 六大能力;Python 绑定(PyO3)目前暴露 `chat/chat_stream/speech`,`image_generate/video_*/embed/transcribe/list_models/list_voices` 待后续版本暴露(见迁移指南 Q1) ## 7. 新 agent 接手指南 @@ -121,36 +153,27 @@ edge-tts / elevenlabs / cartesia / deepgram / assemblyai ### 7.2 验证当前状态 ```bash git checkout feat/aibridge-rust-rewrite -cargo test -p aibridge-core # 期望 810 passed +cargo test -p aibridge-core # 期望 1448 passed cargo test -p aibridge-ffi # 期望 39 passed cargo build --workspace # 0 warning ``` -### 7.3 怎么继续阶段 2b/2c +### 7.3 怎么继续阶段 3 -**模式**:每批 2 个适配器,worktree 并行实现 + 收尾 agent cherry-pick 注册(同阶段 2a)。 +阶段 3 是发布收尾,主要工作: -1. **启动 2 个 worktree agent**(每个实现一个 adapter .rs): - - prompt 要点:读设计文档第 10 节 + Python 对应 `agn/adapters/.py` + `openai_compat.rs`(地基)+ `volcengine_cv.rs`/`more_models.rs Cohere`(独立协议范例);实现 adapter .rs;mockito 单测;**只 git add 自己的 .rs**(不 add mod.rs/factory.rs);commit 返回 hash - - `isolation: "worktree"` + `run_in_background: true` - - worktree 初始可能在 main 分支(无 crates),需 `git checkout feat/aibridge-rust-rewrite` 或基于它建工作分支 -2. **收尾 agent**:cherry-pick 2 个 commit + 注册 mod.rs(pub mod)+ factory.rs(match 分支 + 别名,参考 Python `agn/adapters/factory.py` 的 register)+ 更新 factory 测试 + `cargo test` 全量 + commit -3. **避免 4+ 并行**(会触发 API 限流 429),每批 2 个 +1. **CI 矩阵**:GitHub Actions 平台(linux/macos/windows × amd64/arm64)× 语言。aibridge-ffi 动态库作为构建 artifact 供 Go/JVM/.NET 打包消费。Rust 核心 + 5 绑定各一个 workflow。 +2. **五语言包发布**: + - Python:`maturin build --release` 产 wheel,发布到 PyPI `aibridge` + - Node:`napi build --release` 产 .node,发布到 npm `aibridge` + - Go:提供 libaibridge 安装脚本,Go module `aibridge-go` + - JVM:动态库打进 jar(按平台 classifier),Maven `io.aibridge:aibridge` + - .NET:动态库打进包(runtimes/{rid}/native/),NuGet `AIBridge` +3. **Python 绑定补全**:在 `crates/aibridge-python/src/lib.rs` 的 `#[pymethods] impl Client` 补 `image_generate/video_create/video_poll/embed/transcribe/list_models/list_voices/recommend_voices`,参照已有 `chat`/`speech` 的模式(builder 构造 + RUNTIME.spawn + map_error)。 +4. **文档网站**:mkdocs 或 similar,整合设计文档 + 迁移指南 + 五语言 API。 +5. **v1 归档**:`main` 分支(Python v1)打 `v1.3.3` tag 后归档,README 指向 v2。 -### 7.4 阶段 2b 各适配器要点 -- **anthropic**:Claude messages API,`POST /messages`,header `x-api-key`+`anthropic-version`,流式 SSE 事件(content_block_delta),`ANTHROPIC_MAPPING`。Python: `agn/adapters/anthropic.py` -- **stability**:Stability AI 图像协议,独立。Python: `agn/adapters/stability.py` -- **runway**:视频协议,独立。Python: `agn/adapters/runway.py` -- **pika**:视频协议,独立。Python: `agn/adapters/pika.py` -- **kling**:可灵视频/图像,独立。Python: `agn/adapters/kling.py` - -### 7.5 阶段 2c 音频要点 -- 二进制载荷(TTS 返 audio_data bytes,ASR 接受 file path/URL/bytes/base64) -- edge-tts 免认证(`requires_api_key=false`,加到 `client.rs` 的 `is_free_provider`) -- TTS 音色健康检查/推荐/自动降级(v1.3.3 特性,保留) -- Python: `agn/adapters/audio_adapters.py`(含全部 5 个) - -### 7.6 关键约束(必须遵守) +### 7.4 关键约束(必须遵守) - **中文注释**(项目强制规则,文件模块文档字符串 + 公开项文档注释) - **错误 code 对齐** Rust `aibridge-core/src/error.rs` 的 `code()` 实际值(带 `_error` 后缀) - **mockito 1.x 用法**:`let mut server = mockito::Server::new_async().await;`(不是 `async_try_start`) @@ -158,8 +181,9 @@ cargo build --workspace # 0 warning - **Python/Node 流式**已重构(PyO3 Coroutine / napi async fn),不要再改回 block_on - **Adapter trait** 在 `crates/aibridge-core/src/adapter/base.rs`;工厂注册在 `adapter/factory.rs`(编译期 match,非运行时注册) - **openai_compat.rs** 的 9 个方法已 pub,子适配器组合委托复用 +- **环境变量双前缀兼容**:`AIBRIDGE_*` 新前缀 + `AGN_*` 老前缀并存(见 `config.rs merge_env`) -### 7.7 关键命令 +### 7.5 关键命令 ```bash # Rust 核心 cargo test -p aibridge-core @@ -181,14 +205,52 @@ cd bindings/go && CGO_ENABLED=1 DYLD_LIBRARY_PATH=../../target/debug go run ./ex # JVM 绑定 cd bindings/jvm && ./gradlew run + +# .NET 绑定 +cargo build -p aibridge-ffi +cd bindings/dotnet && dotnet run ``` -## 8. 提交历史(阶段 0-2a) +## 8. 提交历史(阶段 0-3) ``` +(阶段 3) +<待提交> docs: 阶段3 Python v1→v2 迁移指南 + README + 进度更新 + +(阶段 2c 第三批,阶段 2 全部完成) +12925ad feat(aibridge-core): 阶段2c 第三批收尾 注册 deepgram + assemblyai 到工厂(阶段2 全部完成) +3fb46be feat(aibridge-core): 阶段2c assemblyai 适配器 +e533bff feat(aibridge-core): 阶段2c deepgram 适配器 + +(阶段 2c 第二批,TTS) +ac0a3c8 feat(aibridge-core): 阶段2c 第二批收尾 注册 elevenlabs + cartesia 到工厂 +5c2d659 feat(aibridge-core): 阶段2c cartesia 适配器 +0b7c205 feat(aibridge-core): 阶段2c elevenlabs 适配器 + +(阶段 2b+2c 收尾,视频 + 免费 TTS) +03776d2 feat(aibridge-core): 阶段2b+2c 收尾 注册 kling + edge-tts 到工厂(edge-tts 免认证) +d6493c7 feat(aibridge-core): 阶段2c edge-tts 适配器(免费 TTS) +ee89f32 feat(aibridge-core): 阶段2b kling 适配器(可灵) + +(阶段 2b 第二批,视频) +3f05c39 feat(aibridge-core): 阶段2b 第二批收尾 注册 runway + pika 到工厂 +0124c5c feat(aibridge-core): 阶段2b pika 适配器 +0247589 feat(aibridge-core): 阶段2b runway 适配器 + +(阶段 2b 第一批,独立协议) +04fd202 feat(aibridge-core): 阶段2b 第一批收尾 注册 anthropic + stability 到工厂 +e6151da feat(aibridge-core): 阶段2b stability 适配器 +954d153 feat(aibridge-core): 阶段2b anthropic 适配器 + +(阶段 2a 完成) +37a1b9a docs: AIBridge Rust 重构进度文档 + AGENTS.md 接手指引 3880168 feat(aibridge-core): 阶段2a 第三批收尾(emerging_models + chinese,阶段2a 完成) +43e923f feat(aibridge-core): 阶段2a chinese 适配器(中文模型聚合) +ba4a78b feat(aibridge-core): 阶段2a emerging_models 适配器 a0eb96d feat(aibridge-core): 阶段2a 第二批收尾(additional_models + more_models) cb4536a feat(aibridge-core): 阶段2a 第一批收尾(azure + 聚合平台) + +(阶段 0-1) 7b4a79d fix: 阶段1 Python/Node 流式桥接重构 e5df4f3 fix(aibridge-python): dylib 产物名改为 _aibridge 68954d2 feat: 阶段1.5 跨语言一致性测试 + 错误 code 统一 diff --git a/docs/migration-guide.md b/docs/migration-guide.md new file mode 100644 index 0000000..55f49a9 --- /dev/null +++ b/docs/migration-guide.md @@ -0,0 +1,766 @@ +# AIBridge Python v1 → v2 迁移指南 + +> 本指南帮助现有 `agn-sdk`(v1,Python)用户平滑迁移到 `aibridge`(v2,Rust 核心 + 五语言绑定)。 +> v2 是破坏性升级(semver v2.0.0),API 风格有变化,但能力、provider、方法名基本保持一致。 +> 适配对象:已有项目用 `agn-sdk` 的 Python 代码迁移到 `aibridge`。 + +--- + +## 1. 为什么要迁移 + +| 维度 | v1(agn-sdk) | v2(aibridge) | +|---|---|---| +| 实现语言 | Python(~19700 行) | Rust 核心 + 五语言原生绑定 | +| 支持语言 | 仅 Python | Python / JS-TS / Go / JVM / .NET | +| 性能 | Python 原生 | Rust(无 GIL、原生 async、零成本抽象) | +| 类型安全 | Pydantic + `**kwargs` | serde struct + Builder(编译期保证) | +| provider 数 | 38 个 | 38 个(全量迁移,零丢失) | +| 品牌 | agn-sdk | aibridge | + +**迁移收益**:一套 API 五语言通用;Rust 核心更快更稳;显式 struct 替代 `**kwargs`,IDE 补全与编译期检查更好。 + +**迁移成本**:主要是包名、错误类名、参数传递方式的机械替换。方法名与能力基本不变,业务逻辑无需重写。 + +--- + +## 2. 速查表(一页纸变更总览) + +| 变更点 | v1(agn-sdk) | v2(aibridge) | 影响 | +|---|---|---|---| +| 包名 | `agn-sdk` | `aibridge` | `pip install aibridge` | +| 导入 | `from agn import Client` | `from aibridge import Client` | 改 import | +| 错误基类 | `AGNError` | `AibridgeError` | 改异常类名 | +| 错误子类名 | `RateLimitError` 等 | `RateLimitError` 等 | **不变** | +| 错误 code | `RATE_LIMIT_ERROR`(大写) | `rate_limit_error`(snake_case) | 若解析 code 需改 | +| 参数传递 | `**kwargs` + `ChatOptions` | `Request` struct + Builder | 改调用方式 | +| Options 中间层 | `ChatOptions/ImageOptions/...` | 去除,直接 builder | 删掉 Options | +| 流式入口 | `chat(stream=True)` | `chat_stream(req)` 独立方法 | 拆成两个方法 | +| 翻译 | `client.translate(...)` | `transcribe(req, translate=true)` | 合并进 transcribe | +| 方法名 | `chat/image_generate/...` | **不变** | — | +| provider 名 | `openai/agnes/...` | **不变**(别名也保留) | — | +| 环境变量 | `AGN_API_KEY` | `AIBRIDGE_API_KEY`(兼容老 `AGN_*`) | 可不改 | + +--- + +## 3. 包名与导入 + +### v1 + +```python +# 安装:pip install agn-sdk +from agn import Client, Router +from agn import AGNError, RateLimitError, ValidationError +from agn import ChatOptions, ImageOptions, SpeechOptions +``` + +### v2 + +```python +# 安装:pip install aibridge +from aibridge import Client +from aibridge import AibridgeError, RateLimitError, ValidationError +``` + +v2 去除了 `Options` 中间层,不再导出 `ChatOptions/ImageOptions/...`。`Router` 在 Rust 核心已实现,Python 绑定将随版本迭代暴露。 + +--- + +## 4. 错误类对照 + +错误分类完全一致,仅基类改名、code 格式调整。子类名保持不变,便于 `except` 子句平滑迁移。 + +### 4.1 类名对照 + +| v1 | v2 | 说明 | +|---|---|---| +| `AGNError` | `AibridgeError` | 基类(改名) | +| `AuthenticationError` | `AuthenticationError` | 认证失败 | +| `RateLimitError` | `RateLimitError` | 限流 | +| `ValidationError` | `ValidationError` | 参数校验 | +| `ModelNotFoundError` | `ModelNotFoundError` | 模型不存在 | +| `APIError` | `APIError` | Provider API 错误 | +| `NetworkError` | `NetworkError` | 网络错误 | +| `TimeoutError` | `TimeoutError` | 超时 | +| `UnsupportedCapabilityError` | `UnsupportedCapabilityError` | 能力不支持 | +| `ProviderNotFoundError` | `ProviderNotFoundError` | provider 不存在 | +| `VoiceNotAvailableError` | `VoiceNotAvailableError` | 音色不可用 | +| `ServiceUnavailableError` | `ServiceUnavailableError` | 服务暂不可用 | + +### 4.2 code 字段格式变化 + +v1 的 `code` 是大写常量,v2 改为 snake_case(与 Rust 核心对齐): + +```python +# v1 +except RateLimitError as e: + assert e.code == "RATE_LIMIT_ERROR" + +# v2 +except RateLimitError as e: + assert e.code == "rate_limit_error" # snake_case +``` + +若代码里硬编码了大写 code 字符串,迁移时改成 snake_case。异常消息格式 v2 为 `[code] message`。 + +### 4.3 捕获写法迁移 + +```python +# v1 +from agn import AGNError, RateLimitError + +try: + resp = await client.chat(...) +except RateLimitError as e: + print(f"限流,{e.retry_after} 秒后重试") +except AGNError as e: + print(f"其他 SDK 错误: {e}") + +# v2(仅基类名改) +from aibridge import AibridgeError, RateLimitError + +try: + resp = await client.chat(...) +except RateLimitError as e: + print(f"限流,{e.retry_after} 秒后重试") +except AibridgeError as e: + print(f"其他 SDK 错误: {e}") +``` + +--- + +## 5. Client 构造对照 + +### v1 + +```python +client = Client( + provider="agnes", + api_key="your-key", + base_url="https://api.agnes.ai/v1", + timeout=300, + max_retries=3, + retry_delay=2.0, +) +await client.start() +``` + +### v2(Python 绑定) + +```python +client = Client( + provider="agnes", + api_key="your-key", # 关键字参数 + base_url="https://api.agnes.ai/v1", +) +await client.start() +``` + +v2 Python 绑定目前暴露 `api_key` 与 `base_url` 两个关键字参数;`timeout/max_retries/retry_delay` 走环境变量或后续版本暴露。Rust 核心的 `ClientOptions` 完整支持全部连接参数(见下文 Rust 示例)。 + +### v2(Rust 核心,完整参数) + +```rust +use aibridge_core::client::Client; +use aibridge_core::config::ClientOptions; + +let client = Client::new( + "agnes", + ClientOptions::builder() + .api_key("your-key") + .base_url("https://api.agnes.ai/v1") + .timeout(300) + .max_retries(3) + .retry_delay(2.0) + .build(), +)?; +client.start().await?; +``` + +### 免认证 provider + +edge-tts 在 v1/v2 均免认证,构造时不传 `api_key`: + +```python +# v1 +client = Client(provider="edge-tts") + +# v2 +client = Client(provider="edge-tts") +``` + +v2 额外的 `echo` 是 mock 适配器(免认证,用于管线验证与单元测试)。 + +--- + +## 6. 参数传递范式(核心变化) + +这是 v1→v2 最大的变化:**`**kwargs` + `Options` 中间层 → 显式 `Request` struct + Builder 链式调用**。 + +### 6.1 范式对照 + +**v1:三种传参方式混用** + +```python +# 方式 A:独立参数 +resp = await client.chat(model="gpt-4o", messages=[...], temperature=0.7, max_tokens=1000) + +# 方式 B:Options 中间层(options 优先级高于独立参数) +opts = ChatOptions(temperature=0.7, max_tokens=1000, top_p=0.9) +resp = await client.chat(model="gpt-4o", messages=[...], options=opts) + +# 方式 C:**kwargs 透传厂商特有参数 +resp = await client.chat(model="gpt-4o", messages=[...], reasoning_effort="high") +``` + +**v2:统一用 Request builder(Rust 核心)** + +```rust +let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("Hello!")]) + .temperature(0.7) + .max_tokens(1000) + .top_p(0.9) + .extra("reasoning_effort", "high") // 厂商特有参数走 extra + .build(); +let resp = client.chat(req).await?; +``` + +**v2:Python 绑定(已暴露的方法用关键字参数)** + +```python +# Python 绑定对外仍是关键字参数风格,但去掉了 Options 中间层 +resp = await client.chat( + model="gpt-4o", + messages=[{"role": "user", "content": "Hello!"}], + temperature=0.7, + max_tokens=1000, +) +``` + +### 6.2 Options 类全部去除 + +| v1 Options 类 | v2 替代 | +|---|---| +| `ChatOptions` | `ChatRequest::builder(model, messages)` | +| `ImageOptions` | `ImageRequest::builder(model, prompt)` | +| `VideoOptions` | `VideoRequest::builder(model, prompt)` | +| `EmbedOptions` | `EmbedRequest::builder(model, input)` | +| `TranscribeOptions` | `TranscribeRequest::builder(model, file)` | +| `SpeechOptions` | `SpeechRequest::builder(model, input, voice)` | + +`ParameterMapping` 及预置映射常量(`OPENAI_COMPATIBLE_MAPPING` 等)在 v2 是 Rust 适配器内部实现细节,用户不再接触,无需迁移。 + +### 6.3 厂商特有参数透传 + +v1 用 `**kwargs`,v2 用 `extra` 字段(`HashMap`): + +```rust +// v2 Rust:extra 透传 +let req = ChatRequest::builder("gpt-4o", messages) + .extra("reasoning_effort", "high") + .extra("custom_flag", true) + .build(); +``` + +--- + +## 7. 各能力对照(v1 vs v2 示例) + +下表给出六大能力的 v1 与 v2 代码对照。v2 侧同时给出 Rust 核心 API(完整能力)与 Python 绑定 API(已暴露的方法)。echo 适配器的示例可免认证直接运行;真实 provider 示例需替换为有效 API key。 + +### 7.1 文本对话 chat + +**v1(Python)** + +```python +from agn import Client + +client = Client(provider="agnes", api_key="your-key") +await client.start() + +resp = await client.chat( + model="claude-3-opus", + messages=[ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"}, + ], + temperature=0.7, + max_tokens=1000, +) +print(resp.choices[0].message.content) +await client.close() +``` + +**v2(Python 绑定,echo 可直接运行)** + +```python +from aibridge import Client + +client = Client(provider="echo") # 真实用例改 "agnes" + api_key +await client.start() + +resp = await client.chat( + model="echo-chat", + messages=[ + {"role": "system", "content": "You are a helpful assistant."}, + {"role": "user", "content": "Hello!"}, + ], + temperature=0.7, + max_tokens=1000, +) +print(resp.choices[0].message.content) +await client.close() +``` + +**v2(Rust 核心)** + +```rust +use aibridge_core::client::Client; +use aibridge_core::config::ClientOptions; +use aibridge_core::model::chat::{ChatMessage, ChatRequest}; + +let mut client = Client::new("agnes", ClientOptions::builder().api_key("your-key").build())?; +client.start().await?; + +let req = ChatRequest::builder("claude-3-opus", vec![ + ChatMessage::system("You are a helpful assistant."), + ChatMessage::user("Hello!"), +]) +.temperature(0.7) +.max_tokens(1000) +.build(); + +let resp = client.chat(req).await?; +println!("{}", resp.choices[0].message.content.as_deref().unwrap_or("")); +client.close().await?; +``` + +### 7.2 图像生成 image_generate + +**v1(Python)** + +```python +result = await client.image_generate( + model="dall-e-3", + prompt="A beautiful sunset over the ocean", + size="1024x1024", + n=1, +) +print(result.data[0].url) +``` + +**v2(Rust 核心,Python 绑定逐步暴露中)** + +```rust +use aibridge_core::model::image::ImageRequest; + +let req = ImageRequest::builder("dall-e-3", "A beautiful sunset over the ocean") + .size("1024x1024") + .n(1) + .build(); + +let result = client.image_generate(req).await?; +println!("{}", result.data[0].url.as_deref().unwrap_or("")); +``` + +v2 用 `ImageRequest::builder(model, prompt)` 链式构造,替代 v1 的独立参数 + `ImageOptions`。`reference_images` 在 v2 用 `FileInput` 枚举(`Path/Url/Bytes/Base64`)统一表达。 + +### 7.3 视频生成 video_create + video_poll + +**v1(Python)** + +```python +task = await client.video_create( + model="video-gen-1", + prompt="A cat walking through a forest", + width=1280, + height=720, +) +print(task.task_id) + +# 轮询 +status = await client.video_poll(task_id=task.task_id, model="video-gen-1") +print(status.status) +``` + +**v2(Rust 核心)** + +```rust +use aibridge_core::model::video::VideoRequest; + +let req = VideoRequest::builder("video-gen-1", "A cat walking through a forest") + .width(1280) + .height(720) + .build(); + +let task = client.video_create(req).await?; +println!("{}", task.task_id); + +// 轮询 +let status = client.video_poll(&task.task_id, "video-gen-1").await?; +println!("{:?}", status.status); +``` + +`video_poll` 在 v1/v2 签名一致:`(task_id, model)`。`VideoRequest` 的 `mode` 字段用 `VideoMode` 枚举(`text2video/image2video/keyframes/multiimage`),替代 v1 的字符串字面量。 + +### 7.4 文字转语音 speech(TTS) + +**v1(Python)** + +```python +result = await client.speech( + model="tts-1", + input="你好,欢迎使用语音合成", + voice="alloy", + response_format="mp3", + speed=1.0, +) +result.save_to_file("output.mp3") +``` + +**v2(Python 绑定,echo 可直接运行)** + +```python +result = await client.speech( + model="echo-tts", # 真实用例改 "tts-1" + input="hello", + voice="alloy", + response_format="mp3", + speed=1.0, +) +with open("output.mp3", "wb") as f: + f.write(result.audio_data) # bytes +``` + +**v2(Rust 核心)** + +```rust +use aibridge_core::model::audio::SpeechRequest; + +let req = SpeechRequest::builder("tts-1", "你好,欢迎使用语音合成", "alloy") + .response_format("mp3") + .speed(1.0) + .build(); + +let result = client.speech(req).await?; +let audio: Vec = result.audio_data.unwrap_or_default(); +std::fs::write("output.mp3", &audio)?; +``` + +v2 的 `SpeechResult.save_to_file()` 辅助方法在 Python 绑定层尚未暴露,可直接写 `result.audio_data`(bytes)到文件。`voice` 在 v2 Rust 核心用 `VoiceSpec`(支持候选列表自动降级),Python 绑定接受字符串。 + +### 7.5 语音转文字 transcribe(ASR) + +**v1(Python)** + +```python +result = await client.transcribe( + model="whisper-1", + file="/path/to/audio.mp3", + language="zh", + prompt="这是一段关于人工智能的对话", +) +print(result.text) +``` + +**v2(Rust 核心,Python 绑定逐步暴露中)** + +```rust +use aibridge_core::model::audio::TranscribeRequest; +use aibridge_core::model::image::FileInput; + +let req = TranscribeRequest::builder("whisper-1", FileInput::path("/path/to/audio.mp3")) + .language("zh") + .prompt("这是一段关于人工智能的对话") + .build(); + +let result = client.transcribe(req).await?; +println!("{}", result.text); +``` + +v2 用 `FileInput` 枚举统一表达音频输入(`Path/Url/Bytes/Base64`),替代 v1 的 `file` 参数接受多种类型。 + +**翻译(translate)变化**:v1 的 `client.translate(...)` 在 v2 合并进 `transcribe`,用 `translate(true)` 开关: + +```rust +// v2 翻译为英文 +let req = TranscribeRequest::builder("whisper-1", FileInput::path("/path/to/chinese.mp3")) + .translate(true) // 启用翻译模式 + .build(); +let result = client.transcribe(req).await?; // result.task == "translate" +``` + +### 7.6 文本嵌入 embed + +**v1(Python)** + +```python +result = await client.embed( + model="text-embedding-3-small", + input="hello world", +) +print(result.get_embeddings()[0][:5]) +``` + +**v2(Rust 核心,Python 绑定逐步暴露中)** + +```rust +use aibridge_core::model::common::EmbedRequest; + +let req = EmbedRequest::builder("text-embedding-3-small", "hello world").build(); +let result = client.embed(req).await?; +// result.data[0].embedding 为向量 +``` + +--- + +## 8. 流式对话对照 + +### v1:`chat(stream=True)` 一个方法两种返回 + +```python +# v1:stream=False 返 ChatCompletion,stream=True 返 AsyncGenerator +resp = await client.chat(model="gpt-4o", messages=[...], stream=True) +async for chunk in resp: + print(chunk.choices[0].delta.content, end="", flush=True) +``` + +### v2:独立的 `chat_stream` 方法 + +v2 把流式拆成独立方法 `chat_stream`,返回异步迭代器。语义与 v1 一致(逐块产出 `ChatCompletionChunk`)。 + +**v2(Python 绑定,echo 可直接运行)** + +```python +stream = await client.chat_stream( + model="echo-chat", + messages=[{"role": "user", "content": "hello"}], +) +async for chunk in stream: + # chunk.choices[0].delta.content 为增量文本 + print(chunk.choices[0].content, end="", flush=True) +``` + +**v2(Rust 核心)** + +```rust +let req = ChatRequest::builder("gpt-4o", vec![ChatMessage::user("Hello!")]) + .stream(true) + .build(); + +use futures::StreamExt; +let mut stream = client.chat_stream(req).await?; +while let Some(chunk_result) = stream.next().await { + let chunk = chunk_result?; + if let Some(delta) = chunk.choices.first() { + if let Some(content) = &delta.delta.content { + print!("{}", content); + } + } +} +``` + +取消机制:Python 用 `asyncio` 取消协程即可(Rust 侧自动 drop stream → tokio task abort)。 + +--- + +## 9. Provider 列表对照 + +v2 全量迁移 v1 的 38 个 provider,名称与别名完全保留,**迁移零成本**。下表按类别列出。 + +### 9.1 完整 provider 列表 + +| 类别 | provider(主名) | 别名 | +|---|---|---| +| MVP | `openai` `agnes` `volcengine_cv` `gemini` | — | +| 兼容族 | `azure` | — | +| 聚合平台 | `siliconflow` `togetherai` `fireworksai` `cloudflareai` | `sf` / `together` / `fireworks` / `cloudflare`、`workersai` | +| 扩展模型 | `grok` `yi` `sensenova` `hunyuan` `groq` | `xaigrok` / `lingyiwanwu` / `shangtang` / `tencent_hunyuan` | +| 更多模型 | `deepseek` `stepfun` `mistral` `cohere` `perplexity` | `step` | +| 新兴模型 | `ideogram` `luma` `llama` | `ideo` / `dream-machine`、`lumalabs` / `meta-llama`、`meta` | +| 中文模型 | `qwen` `zhipu` `doubao` `ernie` `kimi` `minimax` | — | +| 独立协议 | `anthropic` `stability` `runway` `pika` `kling` | — | +| 音频 TTS | `edge-tts` `elevenlabs` `cartesia` | `edge_tts`、`edge` / `eleven`、`11labs` / `sonic` | +| 音频 ASR | `deepgram` `assemblyai` | `dg` / `assembly`、`aai` | +| Mock | `echo` | — | + +### 9.2 v1 → v2 provider 兼容性 + +- **全部 38 个 provider 保留**:v1 用到的 provider 名在 v2 全部可用,无需改名。 +- **别名保留**:v1 的别名(如 `sf` → `siliconflow`)在 v2 完全一致。 +- **免认证 provider**:`edge-tts` 在 v1/v2 均免认证;v2 新增 `echo` mock 适配器(免认证,测试用)。 +- **v1.3.0 起的"免费 Provider 免认证"特性保留**。 +- **v1.1.0 起的"实时拉取模型列表"(`list_models` 调 provider `/models` 端点)保留**。 +- **v1.3.3 起的"TTS 音色健康检查/推荐/自动降级"保留**。 + +### 9.3 新增 provider + +v2 相对 v1 设计阶段新增 `echo` mock 适配器(阶段 0.6 管线验证用,常驻可用,不调网络)。真实 provider 与 v1.3.3 持平。 + +--- + +## 10. 环境变量对照 + +v2 引入新前缀 `AIBRIDGE_`,同时**兼容老 `AGN_` 前缀**(迁移期并存,平滑过渡)。 + +| 用途 | v1 | v2(新) | v2(兼容老) | +|---|---|---|---| +| 全局 API Key | `AGN_API_KEY` | `AIBRIDGE_API_KEY` | `AGN_API_KEY` | +| Provider 专属 Key | `AGN_OPENAI_API_KEY` | `AIBRIDGE_OPENAI_API_KEY` | `AGN_OPENAI_API_KEY` | +| 全局 Base URL | `AGN_BASE_URL` | `AIBRIDGE_BASE_URL` | `AGN_BASE_URL` | +| Provider 专属 URL | `AGN_OPENAI_BASE_URL` | `AIBRIDGE_OPENAI_BASE_URL` | `AGN_OPENAI_BASE_URL` | +| 轮询 URL(视频) | `AGN_{PROVIDER}_POLL_URL` | `AIBRIDGE_{PROVIDER}_POLL_URL` | `AGN_{PROVIDER}_POLL_URL` | + +**优先级**:代码显式传入 > `AIBRIDGE_{PROVIDER}_*` > `AGN_{PROVIDER}_*` > `AIBRIDGE_API_KEY` > `AGN_API_KEY`。 + +迁移期可继续用老 `AGN_*` 环境变量,无需立即改。建议新项目用 `AIBRIDGE_*` 前缀。 + +--- + +## 11. 破坏性变更清单 + +以下是 v1 用法在 v2 必须修改的项,逐条核对: + +| # | v1 用法 | v2 要求 | 必改 | +|---|---|---|---| +| 1 | `from agn import ...` | `from aibridge import ...` | 是 | +| 2 | `pip install agn-sdk` | `pip install aibridge` | 是 | +| 3 | `except AGNError` | `except AibridgeError` | 是 | +| 4 | `e.code == "RATE_LIMIT_ERROR"` | `e.code == "rate_limit_error"` | 是(若解析 code) | +| 5 | `chat(stream=True)` | `chat_stream(...)` 独立方法 | 是 | +| 6 | `client.translate(...)` | `transcribe(req, translate=true)` | 是 | +| 7 | `ChatOptions(...)` + `options=` | 删除,参数直接传 | 是 | +| 8 | `ImageOptions/VideoOptions/...` | 删除,用 builder | 是 | +| 9 | `**kwargs` 透传厂商参数 | `extra("key", value)` | 是(Rust)/ 关键字参数(Python) | +| 10 | `result.save_to_file()` | 手动写 `audio_data` 到文件 | 是(TTS,Python 绑定暂未暴露辅助方法) | +| 11 | `Client(provider, api_key, base_url, timeout, ...)` 全部位置/关键字参数 | `Client(provider, *, api_key, base_url)`(Python 绑定) | 是(timeout 等走环境变量) | +| 12 | `from agn import ChatOptions, ImageOptions, ...` | 删除(v2 无 Options 类) | 是 | + +**不变项**(确认无需改): +- 方法名:`chat/image_generate/video_create/video_poll/embed/transcribe/speech/list_models/list_voices/recommend_voices` 全部不变 +- provider 名与别名:全部不变 +- 错误子类名:全部不变 +- 异步上下文管理器:`async with client:` 语义不变 +- 响应模型字段名:`ChatCompletion.choices[0].message.content` 等基本对齐 + +--- + +## 12. 迁移检查清单 + +按顺序逐项检查,确保迁移完整: + +- [ ] 依赖:`pip uninstall agn-sdk && pip install aibridge` +- [ ] 全局替换 import:`from agn` → `from aibridge` +- [ ] 错误基类:`AGNError` → `AibridgeError`(子类名不动) +- [ ] 错误 code 字符串:大写 → snake_case(若有硬编码) +- [ ] 删除所有 `ChatOptions/ImageOptions/VideoOptions/EmbedOptions/TranscribeOptions/SpeechOptions` 用法 +- [ ] 流式调用:`chat(stream=True)` → `chat_stream(...)` +- [ ] 翻译调用:`translate(...)` → `transcribe(req, translate=true)` +- [ ] TTS 保存文件:`result.save_to_file(path)` → 手动写 `result.audio_data` +- [ ] Client 构造:`timeout/max_retries/retry_delay` 改走环境变量(或等后续版本暴露) +- [ ] 环境变量:可继续用 `AGN_*`,建议新代码用 `AIBRIDGE_*` +- [ ] provider 名与别名:无需改(全量保留) +- [ ] 运行测试:用 `echo` 适配器做免认证冒烟测试(`Client(provider="echo")`) + +--- + +## 13. 迁移示例:完整脚本对照 + +下面是一个完整脚本的 v1→v2 迁移对照,覆盖 chat + 流式 + speech。 + +### v1 完整脚本 + +```python +import asyncio +from agn import Client, ChatOptions, AGNError, RateLimitError + +async def main(): + client = Client(provider="agnes", api_key="your-key", base_url="https://api.agnes.ai/v1") + async with client: + # 对话 + opts = ChatOptions(temperature=0.7, max_tokens=1000) + resp = await client.chat( + model="claude-3-opus", + messages=[{"role": "user", "content": "Hello!"}], + options=opts, + ) + print(resp.choices[0].message.content) + + # 流式 + stream = await client.chat( + model="claude-3-opus", + messages=[{"role": "user", "content": "讲个笑话"}], + stream=True, + ) + async for chunk in stream: + print(chunk.choices[0].delta.content or "", end="", flush=True) + print() + + # TTS + audio = await client.speech(model="tts-1", input="你好", voice="alloy") + audio.save_to_file("hello.mp3") + +asyncio.run(main()) +``` + +### v2 完整脚本(Python 绑定) + +```python +import asyncio +from aibridge import Client, AibridgeError, RateLimitError + +async def main(): + client = Client(provider="agnes", api_key="your-key", base_url="https://api.agnes.ai/v1") + async with client: + # 对话(去掉 Options,直接传关键字参数) + resp = await client.chat( + model="claude-3-opus", + messages=[{"role": "user", "content": "Hello!"}], + temperature=0.7, + max_tokens=1000, + ) + print(resp.choices[0].message.content) + + # 流式(独立方法 chat_stream) + stream = await client.chat_stream( + model="claude-3-opus", + messages=[{"role": "user", "content": "讲个笑话"}], + ) + async for chunk in stream: + print(chunk.choices[0].content or "", end="", flush=True) + print() + + # TTS(手动写文件) + audio = await client.speech(model="tts-1", input="你好", voice="alloy") + with open("hello.mp3", "wb") as f: + f.write(audio.audio_data) + +asyncio.run(main()) +``` + +--- + +## 14. 常见问题 + +**Q1:v2 Python 绑定为什么部分能力还没暴露?** +A:v2 是 Rust 核心 + 五语言绑定架构。Rust 核心已全部实现 38 provider + 六大能力;Python 绑定(PyO3)目前暴露 `chat/chat_stream/speech`,其余能力(`image_generate/video_*/embed/transcribe/list_models/list_voices`)随版本迭代暴露。急需完整能力的场景可直接用 Rust 核心,或等绑定层补全。 + +**Q2:迁移后性能会有提升吗?** +A:是。Rust 核心无 GIL、原生 async、零成本抽象,IO 密集场景吞吐与延迟显著优于 Python。Python 绑定通过 PyO3 直连 Rust 核心(无 JSON 序列化边界),真实 IO 在 tokio worker 线程执行,不阻塞 asyncio 事件循环。 + +**Q3:v1 和 v2 能并存吗?** +A:能。两者包名不同(`agn-sdk` vs `aibridge`)、import 路径不同(`agn` vs `aibridge`),可在同一环境并存。但建议迁移完成后卸载 v1 避免混淆。 + +**Q4:环境变量必须改吗?** +A:不必。v2 兼容老 `AGN_*` 前缀,迁移期可继续用。新项目建议用 `AIBRIDGE_*`。 + +**Q5:旧版 v1 会归档吗?** +A:会。v2 正式发版后,v1 仓库归档保留,PyPI 上的 `agn-sdk` 不再更新。建议迁移到 `aibridge`。 + +--- + +## 15. 相关文档 + +- [设计文档](superpowers/specs/2026-07-07-aibridge-rust-rewrite-design.md):架构、数据模型、FFI 边界、错误处理 +- [进度文档](PROGRESS.md):当前实施进度与接手指南 +- [README(v2)](../README_aibridge.md):五语言快速开始 + provider 列表 +- [原 README(v1)](../README.md):Python v1 文档(归档参考) From 3eb07ced059346dd9a1f50acf6a56c0cc15e2916 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 04:13:44 +0800 Subject: [PATCH 49/55] =?UTF-8?q?feat:=20=E9=98=B6=E6=AE=B53=20CI=20?= =?UTF-8?q?=E7=9F=A9=E9=98=B5=20+=20=E4=BA=94=E8=AF=AD=E8=A8=80=E6=89=93?= =?UTF-8?q?=E5=8C=85=E8=84=9A=E6=9C=AC?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - .github/workflows/ci.yml: 6 job 跨平台矩阵(Rust core+ffi / Python / Node / Go / JVM / .NET),覆盖 linux/macos/windows × amd64/arm64, 缓存 cargo/npm/gradle,触发于 push 到 feat/aibridge-rust-rewrite + PR - scripts/: 7 个 bash 打包脚本(build-ffi/python/node/go/jvm/dotnet + release 一键构建),中文注释,set -euo pipefail,可本地运行 - package.json: napi.triples 配置多平台目标(defaults + additional) - build.gradle.kts: copyNativeLib task 把 libaibridge 打进 jar (JNA classpath 约定路径 {os}/{arch}/,按平台 classifier) - AIBridge.csproj: NuGet 打包 runtimes/{rid}/native/(.NET 标准) 不实际发布(无 token),仅 CI 配置 + 打包脚本就绪。 --- .github/workflows/ci.yml | 257 +++++++++++++++++++++++ bindings/dotnet/AIBridge/AIBridge.csproj | 16 ++ bindings/jvm/build.gradle.kts | 55 +++++ crates/aibridge-node/package.json | 12 +- scripts/build-dotnet.sh | 54 +++++ scripts/build-ffi.sh | 27 +++ scripts/build-go.sh | 37 ++++ scripts/build-jvm.sh | 25 +++ scripts/build-node.sh | 19 ++ scripts/build-python.sh | 51 +++++ scripts/release.sh | 54 +++++ 11 files changed, 606 insertions(+), 1 deletion(-) create mode 100644 .github/workflows/ci.yml create mode 100755 scripts/build-dotnet.sh create mode 100755 scripts/build-ffi.sh create mode 100755 scripts/build-go.sh create mode 100755 scripts/build-jvm.sh create mode 100755 scripts/build-node.sh create mode 100755 scripts/build-python.sh create mode 100755 scripts/release.sh diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml new file mode 100644 index 0000000..356c2b3 --- /dev/null +++ b/.github/workflows/ci.yml @@ -0,0 +1,257 @@ +# AIBridge CI 矩阵 +# +# 对应设计文档第 12.3 节:平台(linux/macos/windows × amd64/arm64)× 语言。 +# 共 6 个 job:Rust 核心+ffi / Python / Node / Go / JVM / .NET。 +# 触发:push 到 feat/aibridge-rust-rewrite + 所有 PR。 +# +# 说明: +# - 不实际发布(无 token),仅验证跨平台构建 + 产出 artifact 供下载。 +# - cargo 用 Swatinem/rust-cache 缓存;npm 由 setup-node 缓存;gradle 由 setup-gradle 缓存。 +# - ubuntu-24.04-arm 提供 linux arm64 runner;macos-13 提供 macOS amd64 (Intel) runner。 + +name: CI + +on: + push: + branches: [feat/aibridge-rust-rewrite] + pull_request: + +env: + CARGO_TERM_COLOR: always + CARGO_NET_RETRY: 3 + +jobs: + # ────────────────────────────────────────────────────────────────────────── + # 1. Rust 核心 + ffi(跨平台 build + test) + # ────────────────────────────────────────────────────────────────────────── + rust-core-ffi: + name: Rust 核心+ffi (${{ matrix.label }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + include: + - os: ubuntu-latest + label: linux-x64 + - os: ubuntu-24.04-arm + label: linux-arm64 + - os: macos-latest + label: macos-arm64 + - os: macos-13 + label: macos-x64 + - os: windows-latest + label: windows-x64 + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: cargo build(core + ffi) + run: cargo build -p aibridge-core -p aibridge-ffi + - name: cargo test(core + ffi) + run: cargo test -p aibridge-core -p aibridge-ffi + + # ────────────────────────────────────────────────────────────────────────── + # 2. Python 绑定(maturin build wheel) + # macOS 产 universal2 wheel(同时含 amd64+arm64) + # ────────────────────────────────────────────────────────────────────────── + python-bindings: + name: Python wheel (${{ matrix.label }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + include: + - os: ubuntu-latest + label: linux-x64 + - os: ubuntu-24.04-arm + label: linux-arm64 + - os: macos-latest + label: macos-universal2 + - os: windows-latest + label: windows-x64 + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 添加 macOS universal2 交叉编译 target + if: matrix.label == 'macos-universal2' + run: rustup target add x86_64-apple-darwin aarch64-apple-darwin + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: 安装 Python + uses: actions/setup-python@v5 + with: + python-version: '3.12' + - name: 安装 maturin + run: pip install maturin + - name: maturin build wheel + shell: bash + run: | + cd crates/aibridge-python + if [ "${{ matrix.label }}" = "macos-universal2" ]; then + maturin build --release --universal2 + else + maturin build --release + fi + - name: 上传 wheel 产物 + uses: actions/upload-artifact@v4 + with: + name: python-wheel-${{ matrix.label }} + path: crates/aibridge-python/target/wheels/*.whl + if-no-files-found: warn + + # ────────────────────────────────────────────────────────────────────────── + # 3. Node 绑定(napi build .node) + # ────────────────────────────────────────────────────────────────────────── + node-bindings: + name: Node .node (${{ matrix.label }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + include: + - os: ubuntu-latest + label: linux-x64 + - os: ubuntu-24.04-arm + label: linux-arm64 + - os: macos-latest + label: macos-arm64 + - os: macos-13 + label: macos-x64 + - os: windows-latest + label: windows-x64 + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: 安装 Node + uses: actions/setup-node@v4 + with: + node-version: '20' + cache: 'npm' + cache-dependency-path: crates/aibridge-node/package-lock.json + - name: npm install + run: | + cd crates/aibridge-node + npm install + - name: napi build + run: | + cd crates/aibridge-node + npx napi build --platform --release + - name: 上传 .node 产物 + uses: actions/upload-artifact@v4 + with: + name: node-binary-${{ matrix.label }} + path: crates/aibridge-node/*.node + if-no-files-found: warn + + # ────────────────────────────────────────────────────────────────────────── + # 4. Go 绑定(cargo build ffi + go build,CGO) + # windows CGO 标记为 experimental(失败不阻断 CI) + # ────────────────────────────────────────────────────────────────────────── + go-bindings: + name: Go 绑定 (${{ matrix.label }}) + runs-on: ${{ matrix.os }} + continue-on-error: ${{ matrix.experimental }} + strategy: + fail-fast: false + matrix: + include: + - os: ubuntu-latest + label: linux-x64 + experimental: false + - os: macos-latest + label: macos-arm64 + experimental: false + - os: windows-latest + label: windows-x64 + experimental: true + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: 安装 Go + uses: actions/setup-go@v5 + with: + go-version: '1.26' + cache: false + - name: 构建 libaibridge (ffi, release) + run: cargo build -p aibridge-ffi --release + - name: go build + vet (CGO) + shell: bash + env: + CGO_ENABLED: '1' + LIBAIBRIDGE_DIR: ${{ github.workspace }}/target/release + run: | + cd bindings/go + case "$(uname -s)" in + Darwin) export DYLD_LIBRARY_PATH="$LIBAIBRIDGE_DIR" ;; + Linux) export LD_LIBRARY_PATH="$LIBAIBRIDGE_DIR" ;; + esac + go build ./... + go vet ./... + + # ────────────────────────────────────────────────────────────────────────── + # 5. JVM 绑定(cargo build ffi + gradle build,动态库打进 jar) + # ────────────────────────────────────────────────────────────────────────── + jvm-bindings: + name: JVM 绑定 (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, macos-latest, windows-latest] + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: 安装 Java + uses: actions/setup-java@v4 + with: + distribution: temurin + java-version: '21' + - name: 配置 Gradle + uses: gradle/actions/setup-gradle@v4 + - name: 构建 libaibridge (ffi, release) + run: cargo build -p aibridge-ffi --release + - name: gradle build(含 native 进 jar) + shell: bash + run: | + cd bindings/jvm + ./gradlew build -PembedNative=true + + # ────────────────────────────────────────────────────────────────────────── + # 6. .NET 绑定(cargo build ffi + dotnet build,动态库打进 NuGet) + # ────────────────────────────────────────────────────────────────────────── + dotnet-bindings: + name: .NET 绑定 (${{ matrix.os }}) + runs-on: ${{ matrix.os }} + strategy: + fail-fast: false + matrix: + os: [ubuntu-latest, macos-latest, windows-latest] + steps: + - uses: actions/checkout@v4 + - name: 安装 Rust 工具链 + uses: dtolnay/rust-toolchain@stable + - name: 缓存 cargo + uses: Swatinem/rust-cache@v2 + - name: 安装 .NET SDK + uses: actions/setup-dotnet@v4 + with: + dotnet-version: '8.0.x' + - name: 构建 libaibridge (ffi, release) + run: cargo build -p aibridge-ffi --release + - name: dotnet build + shell: bash + run: | + cd bindings/dotnet/AIBridge + dotnet build diff --git a/bindings/dotnet/AIBridge/AIBridge.csproj b/bindings/dotnet/AIBridge/AIBridge.csproj index 9f92aa1..ee30fda 100644 --- a/bindings/dotnet/AIBridge/AIBridge.csproj +++ b/bindings/dotnet/AIBridge/AIBridge.csproj @@ -50,4 +50,20 @@ Link="aibridge.dll" /> + + + AIBridge + true + 2.0.0-alpha.1 + AIBridge .NET 绑定(P/Invoke,直连 aibridge-ffi) + + + + + + diff --git a/bindings/jvm/build.gradle.kts b/bindings/jvm/build.gradle.kts index 5a6998c..adefe52 100644 --- a/bindings/jvm/build.gradle.kts +++ b/bindings/jvm/build.gradle.kts @@ -61,3 +61,58 @@ tasks.withType { tasks.test { useJUnitPlatform() } + +// ────────────────────────────────────────────────────────────────────────── +// 动态库打进 jar(发布分发,设计文档 12.1) +// +// copyNativeLib:把 target/release 的 libaibridge 拷到 build/resources/main/{os}/{arch}/ +// (JNA classpath 约定路径,运行时 Native.load 自动从 jar 内加载)。 +// 由 -PembedNative=true 触发(CI/发布用);本地 ./gradlew run 仍走 java.library.path。 +// jar task 按平台加 classifier(如 darwin-aarch64),便于多平台并行发布。 +// ────────────────────────────────────────────────────────────────────────── + +// 计算 native 库的 OS/arch 标识(JNA 约定:darwin/linux/win32 × x86_64/aarch64) +val nativeOs: String = when { + System.getProperty("os.name").lowercase().contains("mac") -> "darwin" + System.getProperty("os.name").lowercase().contains("linux") -> "linux" + System.getProperty("os.name").lowercase().contains("windows") -> "win32" + else -> "unknown" +} +val nativeArch: String = when (System.getProperty("os.arch").lowercase()) { + "aarch64", "arm64" -> "aarch64" + "x86_64", "amd64" -> "x86_64" + else -> System.getProperty("os.arch") +} + +// 动态库文件名(按平台) +val nativeLibName: String = when (nativeOs) { + "darwin" -> "libaibridge.dylib" + "linux" -> "libaibridge.so" + "win32" -> "aibridge.dll" + else -> "libaibridge.unknown" +} + +// release 动态库目录(cargo build -p aibridge-ffi --release 产物) +val ffiReleaseDir = file("${rootProject.projectDir}/../../target/release") + +// 拷贝 libaibridge 进 jar resources(JNA classpath 约定路径 {os}/{arch}/) +tasks.register("copyNativeLib") { + description = "把 libaibridge 拷进 build/resources/main/{os}/{arch}/(打进 jar)" + group = "build" + from(ffiReleaseDir) { + include(nativeLibName) + } + into(layout.buildDirectory.dir("resources/main/$nativeOs/$nativeArch")) + // 仅 -PembedNative=true 且 release 产物存在时执行(避免本地构建因无产物报错) + onlyIf { + project.hasProperty("embedNative") && file("${ffiReleaseDir}/$nativeLibName").exists() + } +} + +// jar 依赖 copyNativeLib(确保 native 打进 jar),发布时按平台加 classifier +tasks.jar { + dependsOn("copyNativeLib") + if (project.hasProperty("embedNative")) { + archiveClassifier.set("$nativeOs-$nativeArch") + } +} diff --git a/crates/aibridge-node/package.json b/crates/aibridge-node/package.json index 191863c..de8c7c7 100644 --- a/crates/aibridge-node/package.json +++ b/crates/aibridge-node/package.json @@ -24,7 +24,17 @@ }, "napi": { "name": "aibridge", - "triples": {} + "triples": { + "defaults": true, + "additional": [ + "aarch64-apple-darwin", + "aarch64-unknown-linux-gnu", + "aarch64-unknown-linux-musl", + "x86_64-unknown-linux-gnu", + "x86_64-unknown-linux-musl", + "x86_64-pc-windows-msvc" + ] + } }, "scripts": { "build": "napi build --platform", diff --git a/scripts/build-dotnet.sh b/scripts/build-dotnet.sh new file mode 100755 index 0000000..28bab69 --- /dev/null +++ b/scripts/build-dotnet.sh @@ -0,0 +1,54 @@ +#!/usr/bin/env bash +# build-dotnet.sh - 构建 .NET 绑定(dotnet pack,含动态库打进 NuGet) +# +# 产物:bindings/dotnet/AIBridge/bin/Release/*.nupkg +# 动态库放进 runtimes/{rid}/native/(.NET 标准,运行时 NativeLibrary 自动加载)。 +# 若本机无 dotnet SDK,打印提示并退出 0(不阻塞 release.sh)。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +if ! command -v dotnet >/dev/null 2>&1; then + echo "==> 未检测到 dotnet SDK,跳过 .NET 构建" + echo " 安装:https://dotnet.microsoft.com/download" + echo " macOS: brew install --cask dotnet-sdk" + exit 0 +fi + +echo "==> 先构建 libaibridge (release)" +"$REPO_ROOT/scripts/build-ffi.sh" + +# 计算 .NET RID + 动态库文件名 +case "$(uname -s)" in + Darwin) + case "$(uname -m)" in + arm64|aarch64) RID="osx-arm64" ;; + *) RID="osx-x64" ;; + esac + LIB="libaibridge.dylib" ;; + Linux) + case "$(uname -m)" in + aarch64|arm64) RID="linux-arm64" ;; + *) RID="linux-x64" ;; + esac + LIB="libaibridge.so" ;; + MINGW*|MSYS*|CYGWIN*) + RID="win-x64" + LIB="aibridge.dll" ;; + *) + echo "错误:不支持的操作系统 $(uname -s)" >&2 + exit 1 ;; +esac + +LIBDIR="$REPO_ROOT/target/release" +NATIVE_DIR="$REPO_ROOT/bindings/dotnet/AIBridge/runtimes/$RID/native" +mkdir -p "$NATIVE_DIR" +cp "$LIBDIR/$LIB" "$NATIVE_DIR/" +echo "==> 拷贝 $LIB -> runtimes/$RID/native/" + +cd "$REPO_ROOT/bindings/dotnet/AIBridge" +echo "==> dotnet pack(Release)" +dotnet pack -c Release +echo "==> NuGet 产物:" +ls -lh bin/Release/*.nupkg 2>/dev/null || true diff --git a/scripts/build-ffi.sh b/scripts/build-ffi.sh new file mode 100755 index 0000000..46b1955 --- /dev/null +++ b/scripts/build-ffi.sh @@ -0,0 +1,27 @@ +#!/usr/bin/env bash +# build-ffi.sh - 构建 aibridge-ffi 动态库(libaibridge.{so,dylib,dll}) +# +# 产物供 Go / JVM / .NET 绑定消费。 +# 默认 release 模式(产物在 target/release/);可用 PROFILE=debug 切换。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +PROFILE="${PROFILE:-release}" +echo "==> 构建 aibridge-ffi ($PROFILE)" +if [ "$PROFILE" = "release" ]; then + cargo build -p aibridge-ffi --release +else + cargo build -p aibridge-ffi +fi + +# 按平台给出产物文件名 +case "$(uname -s)" in + Darwin) LIB="libaibridge.dylib" ;; + Linux) LIB="libaibridge.so" ;; + MINGW*|MSYS*|CYGWIN*) LIB="aibridge.dll" ;; + *) LIB="libaibridge.(unknown)" ;; +esac + +echo "==> 完成: target/$PROFILE/$LIB" diff --git a/scripts/build-go.sh b/scripts/build-go.sh new file mode 100755 index 0000000..dc84011 --- /dev/null +++ b/scripts/build-go.sh @@ -0,0 +1,37 @@ +#!/usr/bin/env bash +# build-go.sh - 构建 Go 绑定(CGO 调 libaibridge) +# +# 依赖:先产 libaibridge(本脚本自动调用 build-ffi.sh)。 +# 运行 Go 程序需 libaibridge 在动态库搜索路径,见末尾提示(Go 生态惯例)。 +# 可用 RUN_GO_TEST=1 额外跑 go test(部分测试需真实 provider)。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +echo "==> 先构建 libaibridge (release)" +"$REPO_ROOT/scripts/build-ffi.sh" + +LIBDIR="$REPO_ROOT/target/release" +echo "==> 构建 Go 绑定 (CGO)" +cd "$REPO_ROOT/bindings/go" +export CGO_ENABLED=1 +case "$(uname -s)" in + Darwin) export DYLD_LIBRARY_PATH="$LIBDIR:${DYLD_LIBRARY_PATH:-}" ;; + Linux) export LD_LIBRARY_PATH="$LIBDIR:${LD_LIBRARY_PATH:-}" ;; +esac + +go build ./... +echo "==> go build 通过" + +if [ "${RUN_GO_TEST:-0}" = "1" ]; then + go test ./... || echo "(go test 跳过,部分测试需真实 provider)" +fi + +cat < 完成。运行 Go 程序前需让动态库加载器找到 libaibridge: + Linux: export LD_LIBRARY_PATH=$LIBDIR + macOS: export DYLD_LIBRARY_PATH=$LIBDIR + 或将 libaibridge 装到系统库目录(/usr/local/lib)。 + Windows: 把 aibridge.dll 放到可执行文件同目录或 PATH。 +EOF diff --git a/scripts/build-jvm.sh b/scripts/build-jvm.sh new file mode 100755 index 0000000..5f7c30f --- /dev/null +++ b/scripts/build-jvm.sh @@ -0,0 +1,25 @@ +#!/usr/bin/env bash +# build-jvm.sh - 构建 JVM 绑定(gradle build,含动态库打进 jar) +# +# 产物:bindings/jvm/build/libs/aibridge-jvm-*.jar +# 动态库按平台 classifier 打进 jar,放在 JNA classpath 约定路径 {os}/{arch}/ +# (build.gradle.kts 的 copyNativeLib task 负责,-PembedNative=true 触发)。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +echo "==> 先构建 libaibridge (release)" +"$REPO_ROOT/scripts/build-ffi.sh" + +echo "==> gradle build(含 native 进 jar)" +cd "$REPO_ROOT/bindings/jvm" +if [ -x ./gradlew ]; then + ./gradlew build -PembedNative=true --no-daemon +else + # 兜底:本地无 gradle wrapper 时用系统 gradle + gradle build -PembedNative=true +fi + +echo "==> jar 产物:" +ls -lh build/libs/*.jar 2>/dev/null || true diff --git a/scripts/build-node.sh b/scripts/build-node.sh new file mode 100755 index 0000000..98aac44 --- /dev/null +++ b/scripts/build-node.sh @@ -0,0 +1,19 @@ +#!/usr/bin/env bash +# build-node.sh - 构建 Node.js 原生模块(napi-rs .node) +# +# 产物:crates/aibridge-node/aibridge.{platform}-{arch}.node +# 多平台发布:每个平台在 CI 矩阵中分别构建,由 `napi prepublish` 汇总为 +# optionalDependencies 子包(见 package.json 的 napi.triples 配置)。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT/crates/aibridge-node" + +echo "==> npm install" +[ -d node_modules ] || npm install + +echo "==> napi build (release, --platform)" +npx napi build --platform --release + +echo "==> .node 产物:" +ls -lh *.node 2>/dev/null || true diff --git a/scripts/build-python.sh b/scripts/build-python.sh new file mode 100755 index 0000000..f6469e8 --- /dev/null +++ b/scripts/build-python.sh @@ -0,0 +1,51 @@ +#!/usr/bin/env bash +# build-python.sh - 构建 Python wheel(maturin) +# +# 产物:crates/aibridge-python/target/wheels/aibridge-*.whl +# macOS 默认产 universal2 wheel(同时含 amd64+arm64),需 rust target。 +# 可用 BUILD_UNIVERSAL2=0 关闭 universal2。 +set -euo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +cd "$REPO_ROOT" + +# 自动选择 Python 3.10+(pyo3 abi3-py310 要求,系统默认 python3 可能是 3.9) +if [ -z "${PYO3_PYTHON:-}" ]; then + for py in python3.14 python3.13 python3.12 python3.11 python3.10 python3; do + if command -v "$py" >/dev/null 2>&1; then + ver=$("$py" -c 'import sys; print("%d.%d" % sys.version_info[:2])' 2>/dev/null || echo "0.0") + major=${ver%%.*} + minor=${ver#*.} + if [ "${major:-0}" -gt 3 ] || { [ "${major:-0}" -eq 3 ] && [ "${minor:-0}" -ge 10 ]; }; then + export PYO3_PYTHON="$(command -v "$py")" + break + fi + fi + done +fi +if [ -z "${PYO3_PYTHON:-}" ]; then + echo "错误:未找到 Python 3.10+,请安装或设置 PYO3_PYTHON 指向 3.10+ 解释器" >&2 + exit 1 +fi +echo "==> 使用 Python: $PYO3_PYTHON" + +# 确保 maturin 可用 +if ! command -v maturin >/dev/null 2>&1; then + echo "==> 安装 maturin" + "$PYO3_PYTHON" -m pip install -q maturin +fi + +BUILD_ARGS=(build --release) + +# macOS 产 universal2 wheel(amd64+arm64 合一) +if [ "$(uname -s)" = "Darwin" ] && [ "${BUILD_UNIVERSAL2:-1}" = "1" ]; then + echo "==> 添加 universal2 交叉编译 target(x86_64 + aarch64 apple-darwin)" + rustup target add x86_64-apple-darwin aarch64-apple-darwin 2>/dev/null || true + BUILD_ARGS+=(--universal2) +fi + +cd "$REPO_ROOT/crates/aibridge-python" +maturin "${BUILD_ARGS[@]}" + +echo "==> wheel 产物:" +ls -lh target/wheels/*.whl 2>/dev/null || true diff --git a/scripts/release.sh b/scripts/release.sh new file mode 100755 index 0000000..a7b63cd --- /dev/null +++ b/scripts/release.sh @@ -0,0 +1,54 @@ +#!/usr/bin/env bash +# release.sh - 一键构建全部五语言绑定(本地用,不实际发布) +# +# 顺序:ffi(基础)→ python → node → go → jvm → dotnet +# 单语言失败不中断其它语言,最后汇总各语言结果。 +# ffi 失败则整体失败(其余语言依赖 libaibridge)。 +set -uo pipefail + +REPO_ROOT="$(cd "$(dirname "${BASH_SOURCE[0]}")/.." && pwd)" +SCRIPTS="$REPO_ROOT/scripts" + +# 构建某语言并记录结果(不因失败退出) +# 用法:build_one "名称" 脚本路径 +build_one() { + local name="$1" + local script="$2" + echo + echo "====================================================================" + echo "==> 构建 $name" + echo "====================================================================" + if bash "$script"; then + RESULTS+=("$name:OK") + else + RESULTS+=("$name:FAIL") + fi +} + +RESULTS=() +build_one "ffi (libaibridge)" "$SCRIPTS/build-ffi.sh" +build_one "Python wheel" "$SCRIPTS/build-python.sh" +build_one "Node .node" "$SCRIPTS/build-node.sh" +build_one "Go 绑定" "$SCRIPTS/build-go.sh" +build_one "JVM jar" "$SCRIPTS/build-jvm.sh" +build_one ".NET NuGet" "$SCRIPTS/build-dotnet.sh" + +echo +echo "====================================================================" +echo "==> 构建汇总" +echo "====================================================================" +for r in "${RESULTS[@]}"; do + echo " - $r" +done + +# ffi 失败则整体失败 +for r in "${RESULTS[@]}"; do + case "$r" in + "ffi (libaibridge):FAIL") + echo "错误:ffi 构建失败(其它语言依赖它),整体失败" >&2 + exit 1 + ;; + esac +done + +echo "==> 完成(详见上方各语言结果)" From 93ede4d08af6372bc2e9f31d30094bd011215840 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 06:30:37 +0800 Subject: [PATCH 50/55] =?UTF-8?q?feat(aibridge-python):=20=E8=A1=A5?= =?UTF-8?q?=E5=85=A8=E5=85=A8=E9=83=A8=E8=83=BD=E5=8A=9B=E7=BB=91=E5=AE=9A?= =?UTF-8?q?=EF=BC=88image/video/transcribe/embed/list=5Fmodels/list=5Fvoic?= =?UTF-8?q?es=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 在已有 chat/chat_stream/speech 基础上补全 Client 全部能力绑定: 新增方法(Client): - image_generate(model, prompt, size, n, negative_prompt, reference_images, mask, response_format, **kwargs) -> ImageResult - video_create(model, prompt, width, height, num_frames, frame_rate, mode, reference_images, negative_prompt, seed, **kwargs) -> VideoTask - video_poll(task_id, model="") -> VideoStatus - transcribe(model, file, language, prompt, response_format, temperature, **kwargs) -> TranscriptionResult - embed(model, input, **kwargs) -> EmbeddingResult - list_models(model_type=None) -> list[ModelInfo] - list_voices(language=None) -> list[VoiceInfo] - recommend_voices(language=None, gender=None, limit=10) -> list[VoiceInfo] 新增 11 个 #[pyclass] 数据模型: - ImageData / ImageResult - VideoTask / VideoStatus - TranscriptionSegment / TranscriptionWord / TranscriptionResult - EmbeddingItem / EmbeddingResult(含 get_embeddings()) - ModelInfo(type 字段经 #[getter(r#type)] 暴露为关键字 type) - VoiceInfo 实现要点: - **kwargs 经 Option> 接收(async fn 需 'static future), 经 py_to_json 转 serde_json::Value 填入 core extra 透传 - 字符串默认参数用 &str(async &self 方法支持借用,配合字面量默认值) - file/reference_images/mask 支持 str(路径/URL)或 bytes,经 py_to_file_input 转 core FileInput - video_poll 的 model &str 先 to_string 再 spawn(满足 'static) - TaskStatus 经 task_status_to_str 转字符串暴露给 Python - 错误沿用现有 map_error 映射到 AibridgeError 子类层级 验证(echo adapter,免认证): - cargo build -p aibridge-python 通过 - cargo clippy -p aibridge-python -- -D warnings 无警告 - maturin develop 成功 - examples/hello_python_full.py 全部能力通过(chat/image/video_create+poll /transcribe/embed/list_models/list_voices/recommend_voices/speech) - 原 hello_python.py 回归通过 --- crates/aibridge-python/src/lib.rs | 1576 +++++++++++++++++++++++++---- examples/hello_python_full.py | 157 +++ 2 files changed, 1541 insertions(+), 192 deletions(-) create mode 100644 examples/hello_python_full.py diff --git a/crates/aibridge-python/src/lib.rs b/crates/aibridge-python/src/lib.rs index 843e9d0..f910991 100644 --- a/crates/aibridge-python/src/lib.rs +++ b/crates/aibridge-python/src/lib.rs @@ -26,7 +26,7 @@ use std::sync::Arc; use futures::StreamExt; use pyo3::coroutine::Coroutine; use pyo3::prelude::*; -use pyo3::types::{PyBytes, PyString}; +use pyo3::types::{PyBool, PyBytes, PyDict, PyList, PyString, PyTuple}; use tokio::sync::Mutex; use aibridge_core::adapter::ChatStream as CoreChatStream; @@ -35,11 +35,29 @@ use aibridge_core::config::ClientOptions as CoreClientOptions; use aibridge_core::error::AibridgeError as CoreAibridgeError; use aibridge_core::model::audio::{ SpeechRequest as CoreSpeechRequest, SpeechResult as CoreSpeechResult, + TranscribeRequest as CoreTranscribeRequest, TranscriptionResult as CoreTranscriptionResult, + TranscriptionSegment as CoreTranscriptionSegment, TranscriptionWord as CoreTranscriptionWord, }; use aibridge_core::model::chat::{ ChatCompletion as CoreChatCompletion, ChatCompletionChunk as CoreChatCompletionChunk, ChatMessage as CoreChatMessage, ChatRequest as CoreChatRequest, }; +use aibridge_core::model::common::{ + ModelInfo as CoreModelInfo, ModelType as CoreModelType, TaskStatus as CoreTaskStatus, + VideoMode as CoreVideoMode, VoiceInfo as CoreVoiceInfo, +}; +use aibridge_core::model::image::{ + FileInput as CoreFileInput, ImageData as CoreImageData, ImageRequest as CoreImageRequest, + ImageResult as CoreImageResult, +}; +use aibridge_core::model::options::{ + EmbedInput as CoreEmbedInput, EmbedRequest as CoreEmbedRequest, + EmbeddingItem as CoreEmbeddingItem, EmbeddingResult as CoreEmbeddingResult, + EmbeddingVector as CoreEmbeddingVector, +}; +use aibridge_core::model::video::{ + VideoRequest as CoreVideoRequest, VideoStatus as CoreVideoStatus, VideoTask as CoreVideoTask, +}; // =========================================================================== // 全局 tokio runtime @@ -151,6 +169,129 @@ fn map_error(err: CoreAibridgeError) -> PyErr { } } +// =========================================================================== +// Python ↔ core 转换辅助 +// =========================================================================== + +/// 将 Python `**kwargs` 字典转换为 core `extra` 透传参数表 +/// +/// 用于 image_generate / video_create / transcribe / embed 等方法的厂商特有参数透传。 +/// 每个值经 [`py_to_json`] 转为 `serde_json::Value`。 +fn kwargs_to_extra( + d: &Bound<'_, PyDict>, +) -> PyResult> { + let mut map = std::collections::HashMap::new(); + for (k, v) in d.iter() { + let key = k.extract::()?; + map.insert(key, py_to_json(&v)?); + } + Ok(map) +} + +/// 将任意 Python 对象转换为 `serde_json::Value` +/// +/// 支持 None/bool/int/float/str/list/tuple/dict,其余类型兜底为字符串表示。 +/// bool 必须先于 int 判断(Python 中 bool 是 int 子类,`extract::` 对 True 返 Ok)。 +fn py_to_json(obj: &Bound<'_, PyAny>) -> PyResult { + if obj.is_none() { + return Ok(serde_json::Value::Null); + } + // bool 先于 int 判断 + if obj.cast::().is_ok() { + let b: bool = obj.extract()?; + return Ok(serde_json::Value::Bool(b)); + } + if let Ok(i) = obj.extract::() { + return Ok(serde_json::json!(i)); + } + if let Ok(u) = obj.extract::() { + return Ok(serde_json::json!(u)); + } + if let Ok(f) = obj.extract::() { + return Ok(serde_json::json!(f)); + } + if let Ok(s) = obj.extract::() { + return Ok(serde_json::Value::String(s)); + } + if let Ok(d) = obj.cast::() { + let mut map = serde_json::Map::new(); + for (k, v) in d.iter() { + map.insert(k.extract::()?, py_to_json(&v)?); + } + return Ok(serde_json::Value::Object(map)); + } + if let Ok(l) = obj.cast::() { + let mut arr = Vec::new(); + for item in l.iter() { + arr.push(py_to_json(&item)?); + } + return Ok(serde_json::Value::Array(arr)); + } + if let Ok(t) = obj.cast::() { + let mut arr = Vec::new(); + for item in t.iter() { + arr.push(py_to_json(&item)?); + } + return Ok(serde_json::Value::Array(arr)); + } + // 兜底:字符串表示 + Ok(serde_json::Value::String(obj.str()?.to_string())) +} + +/// 将 Python 文件参数(str/bytes)转换为 core `FileInput` +/// +/// - str 以 http(s):// 开头 → `FileInput::Url` +/// - 其他 str → `FileInput::Path` +/// - bytes → `FileInput::Bytes` +fn py_to_file_input(obj: &Bound<'_, PyAny>) -> PyResult { + if let Ok(s) = obj.extract::() { + if s.starts_with("http://") || s.starts_with("https://") { + Ok(CoreFileInput::url(s)) + } else { + Ok(CoreFileInput::path(s)) + } + } else if let Ok(b) = obj.extract::>() { + Ok(CoreFileInput::bytes(b)) + } else { + Err(pyo3::exceptions::PyTypeError::new_err( + "文件参数必须是 str(路径/URL)或 bytes", + )) + } +} + +/// 将 Python embed 输入(str 或 list[str])转换为 core `EmbedInput` +fn py_to_embed_input(obj: &Bound<'_, PyAny>) -> PyResult { + if let Ok(s) = obj.extract::() { + Ok(CoreEmbedInput::Single(s)) + } else if let Ok(v) = obj.extract::>() { + Ok(CoreEmbedInput::Multiple(v)) + } else { + Err(pyo3::exceptions::PyTypeError::new_err( + "input 必须是 str 或 list[str]", + )) + } +} + +/// 将视频生成模式字符串转换为 core `VideoMode` +fn parse_video_mode(s: &str) -> CoreVideoMode { + match s.to_lowercase().as_str() { + "image2video" => CoreVideoMode::Image2Video, + "keyframes" => CoreVideoMode::Keyframes, + "multiimage" => CoreVideoMode::Multiimage, + _ => CoreVideoMode::Text2Video, + } +} + +/// 将 core `TaskStatus` 转为字符串(Python 侧用字符串表示任务状态) +fn task_status_to_str(s: CoreTaskStatus) -> &'static str { + match s { + CoreTaskStatus::Pending => "pending", + CoreTaskStatus::Processing => "processing", + CoreTaskStatus::Success => "success", + CoreTaskStatus::Failed => "failed", + } +} + // =========================================================================== // 数据模型 // =========================================================================== @@ -506,220 +647,882 @@ impl SpeechResult { } } -// =========================================================================== -// 流式迭代器 -// =========================================================================== - -/// 流式对话迭代器 -/// -/// 由 `Client.chat_stream` 返回,实现 `__aiter__`/`__anext__` 协议。 -/// `async for chunk in stream:` 每次取一个 `ChatCompletionChunk`,流结束抛 -/// `StopAsyncIteration`。 -/// -/// 实现说明(不阻塞 asyncio 事件循环): -/// `__anext__` 同步返回一个 PyO3 内置的 [`Coroutine`](可 `await` 的 Python 对象), -/// 其包裹的 Rust future 在全局 tokio runtime 上 `spawn` 消费 core `ChatStream`: -/// - 真实 adapter 的 reqwest IO 在 tokio worker 线程执行,asyncio 线程仅 await -/// `JoinHandle`(Pending 时让出,不阻塞事件循环,其他协程可运行)。 -/// - chunk 就绪后,`Coroutine` 的 `AsyncioWaker` 通过 `asyncio.Future` + -/// `call_soon_threadsafe` 把就绪通知调度回 asyncio 事件循环(PyO3 内置实现, -/// 无需手写 loop 引用),`await` 返回 chunk。 -/// - 流结束:future 返回 `Err(StopAsyncIteration)`;取 chunk 出错:返回对应 -/// `AibridgeError` 子类;正常 chunk:`Ok(chunk)` → `StopIteration(chunk)`。 -/// -/// GIL 处理:future 在 tokio 上 await stream 期间不持 GIL(`spawn` 的 task 在 -/// tokio worker 跑),仅在拿到 chunk 后 `Python::with_gil` 构造 Python 对象。 -#[pyclass] -struct ChatStreamIterator { - /// core 流(None 表示已耗尽) - inner: Arc>>, +/// 图像数据(`ImageResult.data` 元素) +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ImageData { + url: Option, + b64_json: Option, + revised_prompt: Option, } #[pymethods] -impl ChatStreamIterator { - /// 返回自身(异步迭代器协议:`__aiter__` 返回 self) - fn __aiter__(slf: Py) -> Py { - slf +impl ImageData { + #[getter] + fn url(&self) -> Option { + self.url.clone() } - /// 取下一个 chunk(同步返回 `Coroutine`,可 `await`) - /// - /// 返回的 `Coroutine` `await` 后得到 `ChatCompletionChunk`,或抛 - /// `StopAsyncIteration`(流结束)/ 对应 `AibridgeError` 子类(取 chunk 出错)。 - /// - /// 不阻塞事件循环:实际取 chunk 的 future 在 tokio runtime 上推进, - /// `Coroutine` 的 waker 负责把就绪通知桥接回 asyncio 事件循环。 - fn __anext__(&self, py: Python<'_>) -> PyResult> { - let inner = self.inner.clone(); - - // 构造包裹"取下一个 chunk"逻辑的 future。该 future 在 Coroutine 被 - // poll 时推进(poll 发生在 asyncio 线程,持 GIL),但其内部把 stream - // 消费 spawn 到 tokio runtime,await JoinHandle 期间 Pending 让出线程。 - let fut = async move { - // 在 tokio runtime 上消费 stream。spawn 后 await JoinHandle: - // - stream.next()(含真实 reqwest IO)在 tokio worker 线程执行 - // - asyncio 线程仅 poll JoinHandle,Pending 时注册 waker 让出 - let join_result = RUNTIME - .spawn(async move { - let mut guard = inner.lock().await; - match guard.as_mut() { - None => None, - Some(stream) => stream.next().await, - } - }) - .await; - - // JoinError(task panic/取消)→ RuntimeError - let item: Option> = join_result - .map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!( - "chat_stream 消费任务失败: {e}" - )) - })?; + #[getter] + fn b64_json(&self) -> Option { + self.b64_json.clone() + } - // 在 tokio worker 线程拿到 item,需重新进入 GIL 上下文构造 Python 对象。 - // Coroutine future 被 poll 时所在线程(asyncio 线程)已 attached GIL, - // `Python::attach` 在已 attached 线程上直接复用(返回 R),安全构造 pyclass。 - Python::attach(|py: Python<'_>| -> PyResult> { - match item { - // 流结束 → 抛 StopAsyncIteration(await 时终止 async for) - None => Err(pyo3::exceptions::PyStopAsyncIteration::new_err(())), - // 正常 chunk → 返回 chunk 对象(Coroutine 抛 StopIteration(chunk)) - Some(Ok(c)) => { - let chunk = Py::new(py, ChatCompletionChunk::from_core(c))?; - Ok(chunk.into_any()) - } - // 取 chunk 出错 → 抛对应 AibridgeError 子类 - Some(Err(e)) => Err(map_error(e)), - } - }) - }; + #[getter] + fn revised_prompt(&self) -> Option { + self.revised_prompt.clone() + } - // 用 PyO3 内置 Coroutine 包装 future。Coroutine 实现 __await__/__next__/send, - // 可直接被 `await`。其 waker 自动桥接 tokio 唤醒 → asyncio.Future.set_result - // (通过 call_soon_threadsafe),无需手写 asyncio loop 引用。 - let name = PyString::new(py, "ChatStreamIterator.__anext__"); - let coroutine = - pyo3::impl_::coroutine::new_coroutine(&name, Some("ChatStreamIterator"), None, fut); - Py::new(py, coroutine) + fn __repr__(&self) -> String { + format!( + "ImageData(url={:?}, has_b64={})", + self.url, + self.b64_json.is_some() + ) } } -// =========================================================================== -// 客户端 -// =========================================================================== +impl ImageData { + fn from_core(d: CoreImageData) -> Self { + Self { + url: d.url, + b64_json: d.b64_json, + revised_prompt: d.revised_prompt, + } + } +} -/// AIBridge 统一客户端 -/// -/// 对应 Python v1 `Client`,是用户使用 SDK 的唯一入口。 -/// -/// 示例: -/// ```python -/// import asyncio -/// from aibridge import Client -/// -/// async def main(): -/// client = Client(provider="echo") -/// await client.start() -/// resp = await client.chat(model="echo-chat", -/// messages=[{"role": "user", "content": "hello"}]) -/// print(resp.choices[0].message.content) -/// await client.close() -/// -/// asyncio.run(main()) -/// ``` -#[pyclass] -struct Client { - /// core 客户端(用 tokio Mutex 保护以支持 start/close 可变操作) - inner: Arc>, - /// Provider 类型(构造后不变,缓存以避免同步 getter 中 block_on) - provider_type: String, +/// 图像生成结果 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ImageResult { + id: String, + created: u64, + model: String, + data: Vec, } #[pymethods] -impl Client { - /// 创建客户端 - /// - /// 参数: - /// - `provider`: Provider 类型(如 "echo"、"openai") - /// - `api_key`: 可选 API Key(免认证 provider 可省略) - /// - `base_url`: 可选 API Base URL - #[new] - #[pyo3(signature = (provider, *, api_key=None, base_url=None))] - fn new(provider: &str, api_key: Option, base_url: Option) -> PyResult { - let mut opts_builder = CoreClientOptions::builder(); - if let Some(key) = api_key { - opts_builder = opts_builder.api_key(key); - } - if let Some(url) = base_url { - opts_builder = opts_builder.base_url(url); - } - let opts = opts_builder.build(); - let core_client = CoreClient::new(provider, opts).map_err(map_error)?; - let provider_type = core_client.provider_type().to_string(); - Ok(Self { - inner: Arc::new(Mutex::new(core_client)), - provider_type, - }) +impl ImageResult { + #[getter] + fn id(&self) -> String { + self.id.clone() } - /// Provider 类型 #[getter] - fn provider_type(&self) -> String { - self.provider_type.clone() + fn created(&self) -> u64 { + self.created } - /// 启动客户端(初始化适配器) - async fn start(&self) -> PyResult<()> { - let inner = self.inner.clone(); - let result = RUNTIME - .spawn(async move { - let mut client = inner.lock().await; - client.start().await - }) - .await - .map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("start 任务失败: {e}")) - })?; - result.map_err(map_error) + #[getter] + fn model(&self) -> String { + self.model.clone() } - /// 关闭客户端(释放资源) - async fn close(&self) -> PyResult<()> { - let inner = self.inner.clone(); - let result = RUNTIME - .spawn(async move { - let mut client = inner.lock().await; - client.close().await - }) - .await - .map_err(|e| { - pyo3::exceptions::PyRuntimeError::new_err(format!("close 任务失败: {e}")) - })?; - result.map_err(map_error) + #[getter] + fn data(&self) -> Vec { + self.data.clone() } - /// 文本对话 - /// - /// 参数: - /// - `model`: 模型名称 - /// - `messages`: 消息列表(`ChatMessage` 或 `{"role":..., "content":...}` dict) - /// - `temperature`: 可选温度系数 - /// - `max_tokens`: 可选最大 token 数 - #[pyo3(signature = (model, messages, *, temperature=None, max_tokens=None))] - async fn chat( - &self, - model: String, - messages: Vec>, - temperature: Option, - max_tokens: Option, - ) -> PyResult { - // 在持有 GIL 时把 Python 消息转换为 core 消息 - let core_messages = Python::attach(|py| { - messages - .iter() - .map(|m| ChatMessage::to_core(m.bind(py))) - .collect::>>() + fn __repr__(&self) -> String { + format!( + "ImageResult(id={:?}, model={:?}, data_len={})", + self.id, + self.model, + self.data.len() + ) + } +} + +impl ImageResult { + fn from_core(r: CoreImageResult) -> Self { + Self { + id: r.id, + created: r.created, + model: r.model, + data: r.data.into_iter().map(ImageData::from_core).collect(), + } + } +} + +/// 视频任务(创建后返回) +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct VideoTask { + task_id: String, + model: String, + status: String, + created_at: u64, +} + +#[pymethods] +impl VideoTask { + #[getter] + fn task_id(&self) -> String { + self.task_id.clone() + } + + #[getter] + fn model(&self) -> String { + self.model.clone() + } + + #[getter] + fn status(&self) -> String { + self.status.clone() + } + + #[getter] + fn created_at(&self) -> u64 { + self.created_at + } + + fn __repr__(&self) -> String { + format!( + "VideoTask(task_id={:?}, status={:?})", + self.task_id, self.status + ) + } +} + +impl VideoTask { + fn from_core(t: CoreVideoTask) -> Self { + Self { + task_id: t.task_id, + model: t.model, + status: task_status_to_str(t.status).to_string(), + created_at: t.created_at, + } + } +} + +/// 视频任务状态(轮询返回) +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct VideoStatus { + task_id: String, + status: String, + video_url: Option, + progress: Option, + error: Option, + created_at: Option, + updated_at: Option, +} + +#[pymethods] +impl VideoStatus { + #[getter] + fn task_id(&self) -> String { + self.task_id.clone() + } + + #[getter] + fn status(&self) -> String { + self.status.clone() + } + + #[getter] + fn video_url(&self) -> Option { + self.video_url.clone() + } + + #[getter] + fn progress(&self) -> Option { + self.progress + } + + #[getter] + fn error(&self) -> Option { + self.error.clone() + } + + #[getter] + fn created_at(&self) -> Option { + self.created_at + } + + #[getter] + fn updated_at(&self) -> Option { + self.updated_at + } + + fn __repr__(&self) -> String { + format!( + "VideoStatus(task_id={:?}, status={:?}, progress={:?})", + self.task_id, self.status, self.progress + ) + } +} + +impl VideoStatus { + fn from_core(s: CoreVideoStatus) -> Self { + Self { + task_id: s.task_id, + status: task_status_to_str(s.status).to_string(), + video_url: s.video_url, + progress: s.progress, + error: s.error, + created_at: s.created_at, + updated_at: s.updated_at, + } + } +} + +/// 转写词级时间戳 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct TranscriptionWord { + word: String, + start: f64, + end: f64, + confidence: Option, +} + +#[pymethods] +impl TranscriptionWord { + #[getter] + fn word(&self) -> String { + self.word.clone() + } + + #[getter] + fn start(&self) -> f64 { + self.start + } + + #[getter] + fn end(&self) -> f64 { + self.end + } + + #[getter] + fn confidence(&self) -> Option { + self.confidence + } + + fn __repr__(&self) -> String { + format!( + "TranscriptionWord(word={:?}, start={}, end={})", + self.word, self.start, self.end + ) + } +} + +impl TranscriptionWord { + fn from_core(w: CoreTranscriptionWord) -> Self { + Self { + word: w.word, + start: w.start, + end: w.end, + confidence: w.confidence, + } + } +} + +/// 转写分段(带时间戳) +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct TranscriptionSegment { + id: u32, + start: f64, + end: f64, + text: String, + confidence: Option, + speaker: Option, +} + +#[pymethods] +impl TranscriptionSegment { + #[getter] + fn id(&self) -> u32 { + self.id + } + + #[getter] + fn start(&self) -> f64 { + self.start + } + + #[getter] + fn end(&self) -> f64 { + self.end + } + + #[getter] + fn text(&self) -> String { + self.text.clone() + } + + #[getter] + fn confidence(&self) -> Option { + self.confidence + } + + #[getter] + fn speaker(&self) -> Option { + self.speaker.clone() + } + + fn __repr__(&self) -> String { + format!( + "TranscriptionSegment(id={}, text={:?})", + self.id, self.text + ) + } +} + +impl TranscriptionSegment { + fn from_core(s: CoreTranscriptionSegment) -> Self { + Self { + id: s.id, + start: s.start, + end: s.end, + text: s.text, + confidence: s.confidence, + speaker: s.speaker, + } + } +} + +/// 语音转文字结果 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct TranscriptionResult { + text: String, + language: Option, + duration: Option, + task: String, + model: Option, + segments: Option>, + words: Option>, +} + +#[pymethods] +impl TranscriptionResult { + #[getter] + fn text(&self) -> String { + self.text.clone() + } + + #[getter] + fn language(&self) -> Option { + self.language.clone() + } + + #[getter] + fn duration(&self) -> Option { + self.duration + } + + #[getter] + fn task(&self) -> String { + self.task.clone() + } + + #[getter] + fn model(&self) -> Option { + self.model.clone() + } + + #[getter] + fn segments(&self) -> Option> { + self.segments.clone() + } + + #[getter] + fn words(&self) -> Option> { + self.words.clone() + } + + fn __repr__(&self) -> String { + format!( + "TranscriptionResult(text={:?}, language={:?})", + self.text, self.language + ) + } +} + +impl TranscriptionResult { + fn from_core(r: CoreTranscriptionResult) -> Self { + Self { + text: r.text, + language: r.language, + duration: r.duration, + task: r.task, + model: r.model, + segments: r + .segments + .map(|s| s.into_iter().map(TranscriptionSegment::from_core).collect()), + words: r + .words + .map(|w| w.into_iter().map(TranscriptionWord::from_core).collect()), + } + } +} + +/// 单个嵌入项 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct EmbeddingItem { + index: u32, + embedding: Vec, +} + +#[pymethods] +impl EmbeddingItem { + #[getter] + fn index(&self) -> u32 { + self.index + } + + #[getter] + fn embedding(&self) -> Vec { + self.embedding.clone() + } + + fn __repr__(&self) -> String { + format!( + "EmbeddingItem(index={}, dim={})", + self.index, + self.embedding.len() + ) + } +} + +impl EmbeddingItem { + fn from_core(i: CoreEmbeddingItem) -> Self { + let embedding = match i.embedding { + CoreEmbeddingVector::Float(v) => v, + // base64 编码向量未解码,返空(echo 及多数 provider 走 Float) + CoreEmbeddingVector::Base64(_) => Vec::new(), + }; + Self { + index: i.index, + embedding, + } + } +} + +/// 文本嵌入结果 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct EmbeddingResult { + model: String, + data: Vec, + prompt_tokens: Option, + total_tokens: Option, +} + +#[pymethods] +impl EmbeddingResult { + #[getter] + fn model(&self) -> String { + self.model.clone() + } + + #[getter] + fn data(&self) -> Vec { + self.data.clone() + } + + #[getter] + fn prompt_tokens(&self) -> Option { + self.prompt_tokens + } + + #[getter] + fn total_tokens(&self) -> Option { + self.total_tokens + } + + /// 提取所有嵌入向量(便捷访问,等价于 `[item.embedding for item in data]`) + fn get_embeddings(&self) -> Vec> { + self.data.iter().map(|i| i.embedding.clone()).collect() + } + + fn __repr__(&self) -> String { + format!( + "EmbeddingResult(model={:?}, count={})", + self.model, + self.data.len() + ) + } +} + +impl EmbeddingResult { + fn from_core(r: CoreEmbeddingResult) -> Self { + let (prompt_tokens, total_tokens) = r + .usage + .map(|u| (u.prompt_tokens, u.total_tokens)) + .unzip(); + Self { + model: r.model, + data: r.data.into_iter().map(EmbeddingItem::from_core).collect(), + prompt_tokens, + total_tokens, + } + } +} + +/// 模型信息 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct ModelInfo { + id: String, + name: String, + model_type: String, + provider: String, + capabilities: Vec, + max_tokens: Option, + supports_streaming: bool, + description: Option, + created: Option, +} + +#[pymethods] +impl ModelInfo { + #[getter] + fn id(&self) -> String { + self.id.clone() + } + + #[getter] + fn name(&self) -> String { + self.name.clone() + } + + /// 模型类型("chat"/"image"/"video"/"audio") + #[getter(r#type)] + fn model_type(&self) -> String { + self.model_type.clone() + } + + #[getter] + fn provider(&self) -> String { + self.provider.clone() + } + + #[getter] + fn capabilities(&self) -> Vec { + self.capabilities.clone() + } + + #[getter] + fn max_tokens(&self) -> Option { + self.max_tokens + } + + #[getter] + fn supports_streaming(&self) -> bool { + self.supports_streaming + } + + #[getter] + fn description(&self) -> Option { + self.description.clone() + } + + #[getter] + fn created(&self) -> Option { + self.created + } + + fn __repr__(&self) -> String { + format!( + "ModelInfo(id={:?}, type={:?}, provider={:?})", + self.id, self.model_type, self.provider + ) + } +} + +impl ModelInfo { + fn from_core(m: CoreModelInfo) -> Self { + Self { + id: m.id, + name: m.name, + model_type: m.model_type.as_str().to_string(), + provider: m.provider, + capabilities: m.capabilities, + max_tokens: m.max_tokens, + supports_streaming: m.supports_streaming, + description: m.description, + created: m.created, + } + } +} + +/// 音色信息 +#[pyclass(skip_from_py_object)] +#[derive(Debug, Clone)] +struct VoiceInfo { + short_name: Option, + name: Option, + locale: Option, + gender: Option, + voice_id: Option, +} + +#[pymethods] +impl VoiceInfo { + #[getter] + fn short_name(&self) -> Option { + self.short_name.clone() + } + + #[getter] + fn name(&self) -> Option { + self.name.clone() + } + + #[getter] + fn locale(&self) -> Option { + self.locale.clone() + } + + #[getter] + fn gender(&self) -> Option { + self.gender.clone() + } + + #[getter] + fn voice_id(&self) -> Option { + self.voice_id.clone() + } + + fn __repr__(&self) -> String { + format!( + "VoiceInfo(short_name={:?}, locale={:?}, gender={:?})", + self.short_name, self.locale, self.gender + ) + } +} + +impl VoiceInfo { + fn from_core(v: CoreVoiceInfo) -> Self { + Self { + short_name: v.short_name, + name: v.name, + locale: v.locale, + gender: v.gender, + voice_id: v.voice_id, + } + } +} + +// =========================================================================== +// 流式迭代器 +// =========================================================================== + +/// 流式对话迭代器 +/// +/// 由 `Client.chat_stream` 返回,实现 `__aiter__`/`__anext__` 协议。 +/// `async for chunk in stream:` 每次取一个 `ChatCompletionChunk`,流结束抛 +/// `StopAsyncIteration`。 +/// +/// 实现说明(不阻塞 asyncio 事件循环): +/// `__anext__` 同步返回一个 PyO3 内置的 [`Coroutine`](可 `await` 的 Python 对象), +/// 其包裹的 Rust future 在全局 tokio runtime 上 `spawn` 消费 core `ChatStream`: +/// - 真实 adapter 的 reqwest IO 在 tokio worker 线程执行,asyncio 线程仅 await +/// `JoinHandle`(Pending 时让出,不阻塞事件循环,其他协程可运行)。 +/// - chunk 就绪后,`Coroutine` 的 `AsyncioWaker` 通过 `asyncio.Future` + +/// `call_soon_threadsafe` 把就绪通知调度回 asyncio 事件循环(PyO3 内置实现, +/// 无需手写 loop 引用),`await` 返回 chunk。 +/// - 流结束:future 返回 `Err(StopAsyncIteration)`;取 chunk 出错:返回对应 +/// `AibridgeError` 子类;正常 chunk:`Ok(chunk)` → `StopIteration(chunk)`。 +/// +/// GIL 处理:future 在 tokio 上 await stream 期间不持 GIL(`spawn` 的 task 在 +/// tokio worker 跑),仅在拿到 chunk 后 `Python::with_gil` 构造 Python 对象。 +#[pyclass] +struct ChatStreamIterator { + /// core 流(None 表示已耗尽) + inner: Arc>>, +} + +#[pymethods] +impl ChatStreamIterator { + /// 返回自身(异步迭代器协议:`__aiter__` 返回 self) + fn __aiter__(slf: Py) -> Py { + slf + } + + /// 取下一个 chunk(同步返回 `Coroutine`,可 `await`) + /// + /// 返回的 `Coroutine` `await` 后得到 `ChatCompletionChunk`,或抛 + /// `StopAsyncIteration`(流结束)/ 对应 `AibridgeError` 子类(取 chunk 出错)。 + /// + /// 不阻塞事件循环:实际取 chunk 的 future 在 tokio runtime 上推进, + /// `Coroutine` 的 waker 负责把就绪通知桥接回 asyncio 事件循环。 + fn __anext__(&self, py: Python<'_>) -> PyResult> { + let inner = self.inner.clone(); + + // 构造包裹"取下一个 chunk"逻辑的 future。该 future 在 Coroutine 被 + // poll 时推进(poll 发生在 asyncio 线程,持 GIL),但其内部把 stream + // 消费 spawn 到 tokio runtime,await JoinHandle 期间 Pending 让出线程。 + let fut = async move { + // 在 tokio runtime 上消费 stream。spawn 后 await JoinHandle: + // - stream.next()(含真实 reqwest IO)在 tokio worker 线程执行 + // - asyncio 线程仅 poll JoinHandle,Pending 时注册 waker 让出 + let join_result = RUNTIME + .spawn(async move { + let mut guard = inner.lock().await; + match guard.as_mut() { + None => None, + Some(stream) => stream.next().await, + } + }) + .await; + + // JoinError(task panic/取消)→ RuntimeError + let item: Option> = join_result + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!( + "chat_stream 消费任务失败: {e}" + )) + })?; + + // 在 tokio worker 线程拿到 item,需重新进入 GIL 上下文构造 Python 对象。 + // Coroutine future 被 poll 时所在线程(asyncio 线程)已 attached GIL, + // `Python::attach` 在已 attached 线程上直接复用(返回 R),安全构造 pyclass。 + Python::attach(|py: Python<'_>| -> PyResult> { + match item { + // 流结束 → 抛 StopAsyncIteration(await 时终止 async for) + None => Err(pyo3::exceptions::PyStopAsyncIteration::new_err(())), + // 正常 chunk → 返回 chunk 对象(Coroutine 抛 StopIteration(chunk)) + Some(Ok(c)) => { + let chunk = Py::new(py, ChatCompletionChunk::from_core(c))?; + Ok(chunk.into_any()) + } + // 取 chunk 出错 → 抛对应 AibridgeError 子类 + Some(Err(e)) => Err(map_error(e)), + } + }) + }; + + // 用 PyO3 内置 Coroutine 包装 future。Coroutine 实现 __await__/__next__/send, + // 可直接被 `await`。其 waker 自动桥接 tokio 唤醒 → asyncio.Future.set_result + // (通过 call_soon_threadsafe),无需手写 asyncio loop 引用。 + let name = PyString::new(py, "ChatStreamIterator.__anext__"); + let coroutine = + pyo3::impl_::coroutine::new_coroutine(&name, Some("ChatStreamIterator"), None, fut); + Py::new(py, coroutine) + } +} + +// =========================================================================== +// 客户端 +// =========================================================================== + +/// AIBridge 统一客户端 +/// +/// 对应 Python v1 `Client`,是用户使用 SDK 的唯一入口。 +/// +/// 示例: +/// ```python +/// import asyncio +/// from aibridge import Client +/// +/// async def main(): +/// client = Client(provider="echo") +/// await client.start() +/// resp = await client.chat(model="echo-chat", +/// messages=[{"role": "user", "content": "hello"}]) +/// print(resp.choices[0].message.content) +/// await client.close() +/// +/// asyncio.run(main()) +/// ``` +#[pyclass] +struct Client { + /// core 客户端(用 tokio Mutex 保护以支持 start/close 可变操作) + inner: Arc>, + /// Provider 类型(构造后不变,缓存以避免同步 getter 中 block_on) + provider_type: String, +} + +#[pymethods] +impl Client { + /// 创建客户端 + /// + /// 参数: + /// - `provider`: Provider 类型(如 "echo"、"openai") + /// - `api_key`: 可选 API Key(免认证 provider 可省略) + /// - `base_url`: 可选 API Base URL + #[new] + #[pyo3(signature = (provider, *, api_key=None, base_url=None))] + fn new(provider: &str, api_key: Option, base_url: Option) -> PyResult { + let mut opts_builder = CoreClientOptions::builder(); + if let Some(key) = api_key { + opts_builder = opts_builder.api_key(key); + } + if let Some(url) = base_url { + opts_builder = opts_builder.base_url(url); + } + let opts = opts_builder.build(); + let core_client = CoreClient::new(provider, opts).map_err(map_error)?; + let provider_type = core_client.provider_type().to_string(); + Ok(Self { + inner: Arc::new(Mutex::new(core_client)), + provider_type, + }) + } + + /// Provider 类型 + #[getter] + fn provider_type(&self) -> String { + self.provider_type.clone() + } + + /// 启动客户端(初始化适配器) + async fn start(&self) -> PyResult<()> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let mut client = inner.lock().await; + client.start().await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("start 任务失败: {e}")) + })?; + result.map_err(map_error) + } + + /// 关闭客户端(释放资源) + async fn close(&self) -> PyResult<()> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let mut client = inner.lock().await; + client.close().await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("close 任务失败: {e}")) + })?; + result.map_err(map_error) + } + + /// 文本对话 + /// + /// 参数: + /// - `model`: 模型名称 + /// - `messages`: 消息列表(`ChatMessage` 或 `{"role":..., "content":...}` dict) + /// - `temperature`: 可选温度系数 + /// - `max_tokens`: 可选最大 token 数 + #[pyo3(signature = (model, messages, *, temperature=None, max_tokens=None))] + async fn chat( + &self, + model: String, + messages: Vec>, + temperature: Option, + max_tokens: Option, + ) -> PyResult { + // 在持有 GIL 时把 Python 消息转换为 core 消息 + let core_messages = Python::attach(|py| { + messages + .iter() + .map(|m| ChatMessage::to_core(m.bind(py))) + .collect::>>() })?; let mut builder = CoreChatRequest::builder(model, core_messages); @@ -832,6 +1635,384 @@ impl Client { let speech = result.map_err(map_error)?; Ok(SpeechResult::from_core(speech)) } + + /// 图像生成 + /// + /// 参数: + /// - `model`: 图像模型名称 + /// - `prompt`: 提示词 + /// - `size`: 图像尺寸(默认 "1024x1024") + /// - `n`: 生成数量(默认 1) + /// - `negative_prompt`: 负面提示词 + /// - `reference_images`: 参考图列表(图生图),元素为 str(路径/URL)或 bytes + /// - `mask`: 遮罩图(局部重绘),str 或 bytes + /// - `response_format`: 响应格式("url"/"b64_json",默认 "url") + /// - `**kwargs`: 厂商特有参数透传 + #[allow(clippy::too_many_arguments, clippy::type_complexity)] + #[pyo3(signature = (model, prompt, size="1024x1024", n=1, negative_prompt=None, reference_images=None, mask=None, response_format="url", **kwargs))] + async fn image_generate( + &self, + model: String, + prompt: String, + size: &str, + n: u32, + negative_prompt: Option, + reference_images: Option>>, + mask: Option>, + response_format: &str, + kwargs: Option>, + ) -> PyResult { + // 在持 GIL 时转换 Python 对象为 core 类型(**kwargs / reference_images / mask) + let (extra, ref_imgs, mask_input) = Python::attach(|py| -> PyResult<( + std::collections::HashMap, + Vec, + Option, + )> { + let extra = match &kwargs { + Some(d) => kwargs_to_extra(d.bind(py))?, + None => std::collections::HashMap::new(), + }; + let ref_imgs = match &reference_images { + Some(imgs) => imgs + .iter() + .map(|i| py_to_file_input(i.bind(py))) + .collect::>>()?, + None => Vec::new(), + }; + let mask_input = match &mask { + Some(m) => Some(py_to_file_input(m.bind(py))?), + None => None, + }; + Ok((extra, ref_imgs, mask_input)) + })?; + + let mut builder = CoreImageRequest::builder(model, prompt) + .size(size) + .n(n) + .response_format(response_format); + if let Some(np) = negative_prompt { + builder = builder.negative_prompt(np); + } + if !ref_imgs.is_empty() { + builder = builder.reference_images(ref_imgs); + } + if let Some(m) = mask_input { + builder = builder.mask(m); + } + for (k, v) in extra { + builder = builder.extra(k, v); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.image_generate(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("image_generate 任务失败: {e}")) + })?; + + let r = result.map_err(map_error)?; + Ok(ImageResult::from_core(r)) + } + + /// 创建视频生成任务 + /// + /// 参数: + /// - `model`: 视频模型名称 + /// - `prompt`: 提示词 + /// - `width`/`height`: 视频分辨率(默认 1280x720) + /// - `num_frames`: 帧数(部分模型需要) + /// - `frame_rate`: 帧率(默认 24) + /// - `mode`: 生成模式("text2video"/"image2video"/"keyframes"/"multiimage",默认 "text2video") + /// - `reference_images`: 参考图列表(图生视频) + /// - `negative_prompt`: 负面提示词 + /// - `seed`: 随机种子 + /// - `**kwargs`: 厂商特有参数透传 + #[allow(clippy::too_many_arguments)] + #[pyo3(signature = (model, prompt, width=1280, height=720, num_frames=None, frame_rate=24, mode="text2video", reference_images=None, negative_prompt=None, seed=None, **kwargs))] + async fn video_create( + &self, + model: String, + prompt: String, + width: u32, + height: u32, + num_frames: Option, + frame_rate: u32, + mode: &str, + reference_images: Option>>, + negative_prompt: Option, + seed: Option, + kwargs: Option>, + ) -> PyResult { + let (extra, ref_imgs) = Python::attach(|py| -> PyResult<( + std::collections::HashMap, + Vec, + )> { + let extra = match &kwargs { + Some(d) => kwargs_to_extra(d.bind(py))?, + None => std::collections::HashMap::new(), + }; + let ref_imgs = match &reference_images { + Some(imgs) => imgs + .iter() + .map(|i| py_to_file_input(i.bind(py))) + .collect::>>()?, + None => Vec::new(), + }; + Ok((extra, ref_imgs)) + })?; + + let mut builder = CoreVideoRequest::builder(model, prompt) + .width(width) + .height(height) + .frame_rate(frame_rate) + .mode(parse_video_mode(mode)); + if let Some(nf) = num_frames { + builder = builder.num_frames(nf); + } + if !ref_imgs.is_empty() { + builder = builder.reference_images(ref_imgs); + } + if let Some(np) = negative_prompt { + builder = builder.negative_prompt(np); + } + if let Some(s) = seed { + builder = builder.seed(s); + } + for (k, v) in extra { + builder = builder.extra(k, v); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.video_create(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("video_create 任务失败: {e}")) + })?; + + let t = result.map_err(map_error)?; + Ok(VideoTask::from_core(t)) + } + + /// 查询视频任务状态 + /// + /// 参数: + /// - `task_id`: 任务 ID(由 `video_create` 返回) + /// - `model`: 模型名称(部分 Provider 需要,默认空串) + #[pyo3(signature = (task_id, model=""))] + async fn video_poll(&self, task_id: String, model: &str) -> PyResult { + // model 为 &str(借用),spawn 需 'static,故先转 owned + let model = model.to_string(); + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.video_poll(&task_id, &model).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("video_poll 任务失败: {e}")) + })?; + + let s = result.map_err(map_error)?; + Ok(VideoStatus::from_core(s)) + } + + /// 语音转文字 + /// + /// 参数: + /// - `model`: ASR 模型名称 + /// - `file`: 音频文件(str 路径/URL 或 bytes) + /// - `language`: 语言代码(如 "zh"、"en") + /// - `prompt`: 提示词(改善专有名词识别) + /// - `response_format`: 响应格式(默认 "json") + /// - `temperature`: 温度系数(0-1) + /// - `**kwargs`: 厂商特有参数透传 + #[allow(clippy::too_many_arguments)] + #[pyo3(signature = (model, file, language=None, prompt=None, response_format="json", temperature=None, **kwargs))] + async fn transcribe( + &self, + model: String, + file: Py, + language: Option, + prompt: Option, + response_format: &str, + temperature: Option, + kwargs: Option>, + ) -> PyResult { + let (file_input, extra) = Python::attach(|py| -> PyResult<( + CoreFileInput, + std::collections::HashMap, + )> { + let file_input = py_to_file_input(file.bind(py))?; + let extra = match &kwargs { + Some(d) => kwargs_to_extra(d.bind(py))?, + None => std::collections::HashMap::new(), + }; + Ok((file_input, extra)) + })?; + + let mut builder = CoreTranscribeRequest::builder(model, file_input) + .response_format(response_format); + if let Some(l) = language { + builder = builder.language(l); + } + if let Some(p) = prompt { + builder = builder.prompt(p); + } + if let Some(t) = temperature { + builder = builder.temperature(t); + } + for (k, v) in extra { + builder = builder.extra(k, v); + } + let req = builder.build(); + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.transcribe(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("transcribe 任务失败: {e}")) + })?; + + let r = result.map_err(map_error)?; + Ok(TranscriptionResult::from_core(r)) + } + + /// 文本嵌入 + /// + /// 参数: + /// - `model`: 嵌入模型名称 + /// - `input`: 文本(str 或 list[str]) + /// - `**kwargs`: 厂商特有参数透传 + #[pyo3(signature = (model, input, **kwargs))] + async fn embed( + &self, + model: String, + input: Py, + kwargs: Option>, + ) -> PyResult { + let (embed_input, extra) = Python::attach(|py| -> PyResult<( + CoreEmbedInput, + std::collections::HashMap, + )> { + let embed_input = py_to_embed_input(input.bind(py))?; + let extra = match &kwargs { + Some(d) => kwargs_to_extra(d.bind(py))?, + None => std::collections::HashMap::new(), + }; + Ok((embed_input, extra)) + })?; + + let req = CoreEmbedRequest { + model, + input: embed_input, + dimensions: None, + encoding_format: None, + user: None, + extra, + }; + + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.embed(req).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("embed 任务失败: {e}")) + })?; + + let r = result.map_err(map_error)?; + Ok(EmbeddingResult::from_core(r)) + } + + /// 获取可用模型列表 + /// + /// 参数: + /// - `model_type`: 可选模型类型过滤("chat"/"image"/"video"/"audio") + #[pyo3(signature = (model_type=None))] + async fn list_models(&self, model_type: Option) -> PyResult> { + let filter = model_type.map(|s| CoreModelType::from(s.as_str())); + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.list_models(filter).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("list_models 任务失败: {e}")) + })?; + + let models = result.map_err(map_error)?; + Ok(models.into_iter().map(ModelInfo::from_core).collect()) + } + + /// 获取 Provider 可用音色列表 + /// + /// 参数: + /// - `language`: 可选语言过滤(如 "zh-CN") + #[pyo3(signature = (language=None))] + async fn list_voices(&self, language: Option) -> PyResult> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client.list_voices(language.as_deref()).await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("list_voices 任务失败: {e}")) + })?; + + let voices = result.map_err(map_error)?; + Ok(voices.into_iter().map(VoiceInfo::from_core).collect()) + } + + /// 推荐可用音色(按语言/性别过滤) + /// + /// 参数: + /// - `language`: 可选语言过滤(如 "zh-CN") + /// - `gender`: 可选性别过滤("Female"/"Male") + /// - `limit`: 返回数量上限(默认 10) + #[pyo3(signature = (language=None, gender=None, limit=10))] + async fn recommend_voices( + &self, + language: Option, + gender: Option, + limit: usize, + ) -> PyResult> { + let inner = self.inner.clone(); + let result = RUNTIME + .spawn(async move { + let client = inner.lock().await; + client + .recommend_voices(language.as_deref(), gender.as_deref(), limit) + .await + }) + .await + .map_err(|e| { + pyo3::exceptions::PyRuntimeError::new_err(format!("recommend_voices 任务失败: {e}")) + })?; + + let voices = result.map_err(map_error)?; + Ok(voices.into_iter().map(VoiceInfo::from_core).collect()) + } } // =========================================================================== @@ -886,6 +2067,17 @@ fn _aibridge(_py: Python<'_>, m: &Bound<'_, PyModule>) -> PyResult<()> { m.add_class::()?; m.add_class::()?; m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; + m.add_class::()?; // 客户端与流式 m.add_class::()?; diff --git a/examples/hello_python_full.py b/examples/hello_python_full.py new file mode 100644 index 0000000..c04e868 --- /dev/null +++ b/examples/hello_python_full.py @@ -0,0 +1,157 @@ +"""AIBridge Python 绑定全能力验证 + +用 echo mock 适配器(免认证、无网络)端到端验证全部能力: +chat / chat_stream / speech(已有)+ image / video / transcribe / embed +/ list_models / list_voices / recommend_voices(本次补全)。 + +每个能力打印关键字段并断言 echo 固定响应,确保 PyO3 → core 数据流正确。 +""" + +import asyncio + +import aibridge +from aibridge import Client + + +async def main() -> None: + print(f"aibridge 版本: {aibridge.__version__}") + print("=" * 60) + + client = Client(provider="echo") + await client.start() + print(f"已创建客户端,provider_type={client.provider_type}") + + # --- chat(已有,回归验证)--- + print("-" * 60) + resp = await client.chat( + model="echo-chat", + messages=[{"role": "user", "content": "hello"}], + ) + content = resp.choices[0].message.content + print(f"[chat] content = {content!r}") + assert content == "hello [echo]", f"期望 'hello [echo]',实际 {content!r}" + + # --- image_generate --- + print("-" * 60) + img = await client.image_generate(model="echo-image", prompt="cat") + print(f"[image] id={img.id!r} model={img.model!r} data_len={len(img.data)}") + print(f"[image] b64_json 非空: {img.data[0].b64_json is not None}") + print(f"[image] revised_prompt={img.data[0].revised_prompt!r}") + assert len(img.data) == 1, f"期望 1 张图,实际 {len(img.data)}" + assert img.data[0].b64_json is not None, "echo 应返回 b64_json" + assert img.data[0].revised_prompt == "cat", "revised_prompt 应回显 prompt" + + # --- video_create + video_poll --- + print("-" * 60) + task = await client.video_create(model="echo-video", prompt="cat walking") + print(f"[video] task_id={task.task_id!r} status={task.status!r}") + assert task.task_id == "echo-task-1", f"期望 'echo-task-1',实际 {task.task_id!r}" + assert task.status == "success", f"期望 'success',实际 {task.status!r}" + + status = await client.video_poll(task.task_id) + print( + f"[video] poll: status={status.status!r} " + f"video_url={status.video_url!r} progress={status.progress}" + ) + assert status.status == "success" + assert status.video_url == "https://example.com/echo.mp4" + assert status.progress == 100 + + # --- transcribe --- + print("-" * 60) + tr = await client.transcribe(model="echo-asr", file="audio.mp3") + print( + f"[asr] text={tr.text!r} language={tr.language!r} " + f"duration={tr.duration} task={tr.task!r}" + ) + assert tr.text == "echo transcription", f"期望 'echo transcription',实际 {tr.text!r}" + assert tr.language == "zh" + assert tr.task == "transcribe" + + # --- embed --- + print("-" * 60) + emb = await client.embed(model="echo-embed", input=["hello"]) + print( + f"[embed] model={emb.model!r} count={len(emb.data)} " + f"prompt_tokens={emb.prompt_tokens} total_tokens={emb.total_tokens}" + ) + vectors = emb.get_embeddings() + print(f"[embed] get_embeddings() = {vectors}") + assert len(emb.data) == 1, f"期望 1 个向量,实际 {len(emb.data)}" + assert emb.data[0].index == 0 + assert len(vectors) == 1 and len(vectors[0]) == 3, "echo 应返 1 个 3 维向量" + assert emb.prompt_tokens == 1 and emb.total_tokens == 1 + + # embed 单条字符串输入 + emb_single = await client.embed(model="echo-embed", input="hello") + assert len(emb_single.data) == 1, "单条字符串输入应返 1 个向量" + + # --- list_models --- + print("-" * 60) + models = await client.list_models() + model_ids = [m.id for m in models] + print(f"[models] 共 {len(models)} 个: {model_ids}") + assert len(models) == 6, f"期望 6 个 echo 模型,实际 {len(models)}" + assert "echo-chat" in model_ids and "echo-image" in model_ids + # 验证 ModelInfo 字段 + chat_model = next(m for m in models if m.id == "echo-chat") + print( + f"[models] echo-chat: type={chat_model.type!r} provider={chat_model.provider!r} " + f"supports_streaming={chat_model.supports_streaming} capabilities={chat_model.capabilities}" + ) + assert chat_model.type == "chat" + assert chat_model.provider == "echo" + assert chat_model.supports_streaming is True + + # 按类型过滤 + images = await client.list_models(model_type="image") + print(f"[models] image 过滤: {[m.id for m in images]}") + assert len(images) == 1 and images[0].id == "echo-image" + + # --- list_voices --- + print("-" * 60) + voices = await client.list_voices() + print( + f"[voices] 共 {len(voices)} 个: " + f"{[(v.short_name, v.locale, v.gender) for v in voices]}" + ) + assert len(voices) == 2, f"期望 2 个音色,实际 {len(voices)}" + assert voices[0].short_name == "echo-voice-1" + assert voices[0].locale == "zh-CN" + assert voices[0].gender == "Female" + + # --- recommend_voices --- + print("-" * 60) + # echo 的 list_voices 忽略 language 参数(返全部),默认 recommend_voices 仅按 gender 过滤; + # 故此处验证调用成功 + limit 转发 + gender 过滤(证明绑定正确转发参数) + all_voices = await client.recommend_voices() + print(f"[recommend] 全部: {[(v.short_name, v.locale) for v in all_voices]}") + assert len(all_voices) == 2, f"期望 2 个,实际 {len(all_voices)}" + + limited = await client.recommend_voices(limit=1) + print(f"[recommend] limit=1: {[v.short_name for v in limited]}") + assert len(limited) == 1, "limit=1 应只返 1 个" + + female = await client.recommend_voices(gender="Female") + print(f"[recommend] gender=Female: {[(v.short_name, v.gender) for v in female]}") + assert len(female) == 1 and female[0].gender == "Female" + + # --- speech(已有,回归验证)--- + print("-" * 60) + result = await client.speech(model="echo-tts", input="hello", voice="alloy") + audio = result.audio_data + print(f"[speech] len(audio_data)={len(audio)} format={result.format!r}") + assert len(audio) == 15, f"期望 15 字节,实际 {len(audio)}" + + # --- 关闭 --- + print("-" * 60) + await client.close() + print("[close] 客户端已关闭") + + print("=" * 60) + print("全部通过:chat / image / video / transcribe / embed / " + "list_models / list_voices / recommend_voices / speech") + + +if __name__ == "__main__": + asyncio.run(main()) From 3d4b83a4da61d737ddc4b9651e13157fa04c1a10 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Wed, 8 Jul 2026 07:04:55 +0800 Subject: [PATCH 51/55] =?UTF-8?q?test:=20=E9=98=B6=E6=AE=B53=20=E7=9C=9F?= =?UTF-8?q?=E5=AE=9E=20API=20=E6=B5=8B=E8=AF=95=E8=84=9A=E6=9C=AC=20+=20.e?= =?UTF-8?q?nv.example=20=E6=A8=A1=E6=9D=BF?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - .env.example:四 provider(agnes/openai/gemini/火山)key 模板,可上传 GitHub - examples/test_real_providers.py:从 .env 读 key,list_models() 实时拉真实 model 名,测四 provider 各能力,不硬编码 key - .env 已 gitignore,不上传 GitHub --- .env.example | 40 +++++--- examples/test_real_providers.py | 173 ++++++++++++++++++++++++++++++++ 2 files changed, 197 insertions(+), 16 deletions(-) create mode 100644 examples/test_real_providers.py diff --git a/.env.example b/.env.example index 717402d..ad9fb99 100644 --- a/.env.example +++ b/.env.example @@ -1,20 +1,28 @@ -# Agnes AI -AGN_AGNES_API_KEY=your-agnes-api-key -AGN_AGNES_BASE_URL=https://api.agnes.ai/v1 +# AIBridge 真实 API 测试环境变量 +# +# ⚠️ 重要:本文件是模板(.env.example),可上传 GitHub。 +# 实际 API key 请填在 .env 文件里(已 gitignore,不会上传 GitHub)。 +# +# 用法: +# 1. cp .env.example .env +# 2. 在 .env 填入真实 API key(等号右边) +# 3. python examples/test_real_providers.py +# +# 测试脚本会从 .env 读 key,并用 list_models() 实时拉取真实 model 名字(model 经常变,不硬编码)。 -# OpenAI -AGN_OPENAI_API_KEY=your-openai-api-key +# Agnes AI(https://agnes.ai) +AGNES_API_KEY= +AGNES_BASE_URL=https://api.agnes.ai/v1 -# Azure OpenAI -AGN_AZURE_API_KEY=your-azure-api-key -AGN_AZURE_BASE_URL=https://your-resource.openai.azure.com/openai/deployments/your-deployment -AGN_AZURE_API_VERSION=2024-02-15-preview +# OpenAI(https://platform.openai.com) +OPENAI_API_KEY= +OPENAI_BASE_URL=https://api.openai.com/v1 -# Runway -AGN_RUNWAY_API_KEY=your-runway-api-key +# Google Gemini(https://ai.google.dev) +GEMINI_API_KEY= +GEMINI_BASE_URL=https://generativelanguage.googleapis.com/v1beta -# Pika -AGN_PIKA_API_KEY=your-pika-api-key - -# Stability AI -AGN_STABILITY_API_KEY=your-stability-api-key +# 火山引擎 CV / 方舟(https://www.volcengine.com) +# base_url 从方舟控制台获取(含 endpoint ID,形如 https://ark.cn-beijing.volces.com/api/v3) +VOLCENGINE_CV_API_KEY= +VOLCENGINE_CV_BASE_URL= diff --git a/examples/test_real_providers.py b/examples/test_real_providers.py new file mode 100644 index 0000000..4081c4d --- /dev/null +++ b/examples/test_real_providers.py @@ -0,0 +1,173 @@ +""" +AIBridge 真实 API 测试脚本 + +从 .env 读 API key,用 list_models() 实时拉取真实 model 名字(model 经常变,不硬编码), +测试四 provider(agnes/openai/gemini/火山)的各能力。 + +⚠️ 本脚本不含 API key(从 .env 读),可上传 GitHub。 + .env 含 key,已在 .gitignore,不上传 GitHub。 + +用法: +1. cp .env.example .env,在 .env 填入真实 API key +2. maturin develop -m crates/aibridge-python/Cargo.toml(装 aibridge 到 Python) +3. python examples/test_real_providers.py +""" + +import asyncio +import os +import sys +from aibridge import Client + + +def load_env(path: str = ".env") -> dict: + """读取 .env 文件(不依赖 python-dotenv,手动 parse)""" + env = {} + if not os.path.exists(path): + return env + with open(path, encoding="utf-8") as f: + for line in f: + line = line.strip() + if not line or line.startswith("#"): + continue + if "=" in line: + k, v = line.split("=", 1) + env[k.strip()] = v.strip() + return env + + +# 四 provider 配置:provider 名 / key 环境变量 / base_url 环境变量 / 能力 +PROVIDERS = [ + { + "name": "agnes", + "key": "AGNES_API_KEY", + "base_url": "AGNES_BASE_URL", + "caps": ["chat", "image", "video", "embed"], + }, + { + "name": "openai", + "key": "OPENAI_API_KEY", + "base_url": "OPENAI_BASE_URL", + "caps": ["chat", "image", "embed"], + }, + { + "name": "gemini", + "key": "GEMINI_API_KEY", + "base_url": "GEMINI_BASE_URL", + "caps": ["chat", "image", "embed"], + }, + { + "name": "volcengine_cv", + "key": "VOLCENGINE_CV_API_KEY", + "base_url": "VOLCENGINE_CV_BASE_URL", + "caps": ["image", "video"], + }, +] + + +def pick_model(models, model_type: str): + """从 list_models 结果里按类型挑第一个 model(type 大小写容错)""" + for m in models: + try: + if str(m.type).lower() == model_type: + return m + except Exception: + continue + return None + + +async def test_provider(cfg: dict, env: dict) -> None: + name = cfg["name"] + api_key = env.get(cfg["key"], "") + if not api_key: + print(f"\n=== {name} === 跳过(未填 {cfg['key']})") + return + base_url = env.get(cfg["base_url"], "") + kwargs = {"provider": name, "api_key": api_key} + if base_url: + kwargs["base_url"] = base_url + print(f"\n=== {name} === (base_url={base_url or '默认'})") + try: + client = Client(**kwargs) + await client.start() + except Exception as e: + print(f" 创建/启动失败: {e}") + return + + try: + # 1. list_models 实时拉取真实 model 名字 + try: + models = await client.list_models() + print(f" list_models: {len(models)} 个模型") + for m in models[:5]: + print(f" - {m.id} ({m.type})") + if len(models) > 5: + print(f" ... 共 {len(models)} 个") + except Exception as e: + print(f" list_models 失败: {e}") + models = [] + + chat_m = pick_model(models, "chat") + image_m = pick_model(models, "image") + video_m = pick_model(models, "video") + embed_m = pick_model(models, "embed") + + # 2. chat + if "chat" in cfg["caps"] and chat_m: + try: + r = await client.chat( + model=chat_m.id, + messages=[{"role": "user", "content": "说一个字"}], + ) + content = r.choices[0].message.content if r.choices else "(空)" + print(f" chat({chat_m.id}): {str(content)[:30]}") + except Exception as e: + print(f" chat({chat_m.id}) 失败: {e}") + + # 3. image + if "image" in cfg["caps"] and image_m: + try: + img = await client.image_generate( + model=image_m.id, prompt="a cute cat" + ) + ok = bool(img.data) + print(f" image({image_m.id}): {'OK' if ok else 'FAIL'}") + except Exception as e: + print(f" image({image_m.id}) 失败: {e}") + + # 4. video(创建任务,不轮询,避免长时间等待) + if "video" in cfg["caps"] and video_m: + try: + task = await client.video_create( + model=video_m.id, prompt="a cat walking" + ) + print(f" video({video_m.id}): task_id={task.task_id}") + except Exception as e: + print(f" video({video_m.id}) 失败: {e}") + + # 5. embed + if "embed" in cfg["caps"] and embed_m: + try: + e = await client.embed(model=embed_m.id, input=["hello"]) + vecs = e.get_embeddings() + dim = len(vecs[0]) if vecs else 0 + print(f" embed({embed_m.id}): {len(vecs)} 条, {dim} 维") + except Exception as e: + print(f" embed({embed_m.id}) 失败: {e}") + finally: + await client.close() + + +async def main() -> None: + env = load_env() + if not env: + print("未找到 .env,请先 cp .env.example .env 并填入 API key") + sys.exit(1) + print("AIBridge 真实 API 测试") + print("(model 名由 list_models() 实时拉取,不硬编码)") + for cfg in PROVIDERS: + await test_provider(cfg, env) + print("\n测试完成。如有失败项,把错误信息发给我修复。") + + +if __name__ == "__main__": + asyncio.run(main()) From aad7ff33579cad8d583f2b3e377285cf26c2c776 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Thu, 9 Jul 2026 07:34:49 +0800 Subject: [PATCH 52/55] =?UTF-8?q?fix(aibridge-agnes):=20video=20mode=20?= =?UTF-8?q?=E6=98=A0=E5=B0=84=E5=AF=B9=E9=BD=90=E6=96=B9=E8=88=9F=20agnes?= =?UTF-8?q?=20=E5=B9=B3=E5=8F=B0=EF=BC=88ti2vid/keyframes/multi=5Freferenc?= =?UTF-8?q?e=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit 真实 API 测试发现 video_create 报 [validation_error] Input should be 'ti2vid', 'keyframes' or 'multi_reference'。原实现直接把统一 VideoMode 序列化为 text2video/image2video/keyframes/multiimage 透传,与 agnes 平台 要求的 mode 不一致。 新增 map_video_mode 将统一 VideoMode 显式映射到 agnes 平台 mode: - Text2Video / Image2Video -> ti2vid - Keyframes -> keyframes - Multiimage -> multi_reference build_video_body 改用映射函数;更新 5 处单测断言;新增 map_video_mode 单测覆盖全部 4 种 VideoMode 映射。 --- crates/aibridge-core/src/adapters/agnes.rs | 61 +++++++++++++++++++--- 1 file changed, 54 insertions(+), 7 deletions(-) diff --git a/crates/aibridge-core/src/adapters/agnes.rs b/crates/aibridge-core/src/adapters/agnes.rs index bb280a0..749f411 100644 --- a/crates/aibridge-core/src/adapters/agnes.rs +++ b/crates/aibridge-core/src/adapters/agnes.rs @@ -204,6 +204,31 @@ impl AgnesAdapter { // ==================== 视频协议(Agnes 特有) ==================== + /// 将统一 `VideoMode` 映射为 Agnes 平台要求的 mode 字符串 + /// + /// Agnes Video V2.0 平台只接受三种 mode(真实 API 校验,否则报 + /// `[validation_error] Input should be 'ti2vid', 'keyframes' or + /// 'multi_reference'`): + /// - `ti2vid`:文本/图像生成视频(文生视频 + 单图生视频) + /// - `keyframes`:关键帧模式(首尾帧) + /// - `multi_reference`:多参考图模式 + /// + /// 统一 `VideoMode` 到 Agnes mode 的映射: + /// - `Text2Video` / `Image2Video` -> `ti2vid` + /// - `Keyframes` -> `keyframes` + /// - `Multiimage` -> `multi_reference` + /// + /// 注意:Python v1 `agnes.py` 直接透传 mode(同样存在该 bug),此处以 + /// Agnes 平台实际校验为准做显式映射。 + fn map_video_mode(mode: crate::model::common::VideoMode) -> &'static str { + use crate::model::common::VideoMode; + match mode { + VideoMode::Text2Video | VideoMode::Image2Video => "ti2vid", + VideoMode::Keyframes => "keyframes", + VideoMode::Multiimage => "multi_reference", + } + } + /// 构造视频创建请求体 /// /// 对应 Python v1 `agnes.py:video_create` 的 body 构造逻辑: @@ -231,8 +256,8 @@ impl AgnesAdapter { if req.frame_rate != 0 { body["frame_rate"] = json!(req.frame_rate); } - // mode 序列化为小写字符串(text2video / image2video / keyframes / multiimage) - body["mode"] = json!(serde_json::to_value(req.mode).unwrap_or(json!("text2video"))); + // mode 映射为 Agnes 平台要求的字符串(ti2vid / keyframes / multi_reference) + body["mode"] = json!(Self::map_video_mode(req.mode)); if let Some(seed) = req.seed { body["seed"] = json!(seed); } @@ -743,7 +768,7 @@ mod tests { .match_body(mockito::Matcher::PartialJson(json!({ "model": "seedance-2.0", "prompt": "a cat walking", - "mode": "text2video" + "mode": "ti2vid" }))) .with_status(200) .with_body(body.to_string()) @@ -792,7 +817,7 @@ mod tests { "prompt": "a cat", "width": 1920, "height": 1080, - "mode": "image2video", + "mode": "ti2vid", "seed": 42 }))) .with_status(200) @@ -817,7 +842,7 @@ mod tests { let mock = server .mock("POST", "/videos") .match_body(mockito::Matcher::PartialJson(json!({ - "mode": "image2video", + "mode": "ti2vid", "extra_body": {"image": "https://example.com/ref.png"} }))) .with_status(200) @@ -871,7 +896,7 @@ mod tests { let mock = server .mock("POST", "/videos") .match_body(mockito::Matcher::PartialJson(json!({ - "mode": "multiimage", + "mode": "multi_reference", "extra_body": { "image": ["https://example.com/a.png", "https://example.com/b.png"] } @@ -1353,13 +1378,35 @@ mod tests { assert_eq!(s.video_url.as_deref(), Some("https://x.com/nested.mp4")); } + #[test] + fn map_video_mode_aligns_with_agnes_platform() { + // Agnes 平台只接受 ti2vid / keyframes / multi_reference 三种 mode + // 文生/图生视频统一映射为 ti2vid + assert_eq!( + AgnesAdapter::map_video_mode(VideoMode::Text2Video), + "ti2vid" + ); + assert_eq!( + AgnesAdapter::map_video_mode(VideoMode::Image2Video), + "ti2vid" + ); + assert_eq!( + AgnesAdapter::map_video_mode(VideoMode::Keyframes), + "keyframes" + ); + assert_eq!( + AgnesAdapter::map_video_mode(VideoMode::Multiimage), + "multi_reference" + ); + } + #[test] fn build_video_body_text2video_minimal() { let req = VideoRequest::builder("seedance-2.0", "a cat").build(); let body = AgnesAdapter::build_video_body(&req); assert_eq!(body["model"], "seedance-2.0"); assert_eq!(body["prompt"], "a cat"); - assert_eq!(body["mode"], "text2video"); + assert_eq!(body["mode"], "ti2vid"); // 无参考图像时不应有 extra_body assert!(body.get("extra_body").is_none()); } From 88e5ace58c8d82eafdeb2ced0ba11ede880c777d Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Thu, 9 Jul 2026 07:36:08 +0800 Subject: [PATCH 53/55] =?UTF-8?q?test:=20=E7=9C=9F=E5=AE=9E=20API=20?= =?UTF-8?q?=E6=B5=8B=E8=AF=95=E8=84=9A=E6=9C=AC=EF=BC=88agnes/volcengine?= =?UTF-8?q?=20=E5=8D=95=20provider=EF=BC=89+=20.env.example=20=E8=A1=A5=20?= =?UTF-8?q?endpoint=20=E5=8F=98=E9=87=8F?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - agnes 全能力(chat/image/video)真实 API 验证通过 - 火山 list_models 通,image/video 需接入点 ID(方舟规则) - .env.example 加 VOLCENGINE_IMAGE/VIDEO_ENDPOINT 变量 + base_url 说明 --- .env.example | 9 ++++- examples/test_agnes.py | 74 +++++++++++++++++++++++++++++++++++++ examples/test_volcengine.py | 73 ++++++++++++++++++++++++++++++++++++ 3 files changed, 154 insertions(+), 2 deletions(-) create mode 100644 examples/test_agnes.py create mode 100644 examples/test_volcengine.py diff --git a/.env.example b/.env.example index ad9fb99..8c9b087 100644 --- a/.env.example +++ b/.env.example @@ -23,6 +23,11 @@ GEMINI_API_KEY= GEMINI_BASE_URL=https://generativelanguage.googleapis.com/v1beta # 火山引擎 CV / 方舟(https://www.volcengine.com) -# base_url 从方舟控制台获取(含 endpoint ID,形如 https://ark.cn-beijing.volces.com/api/v3) +# base_url 固定 https://ark.cn-beijing.volces.com/api/v3(agent plan/coding plan/普通 api 都是订阅计划,地址统一) VOLCENGINE_CV_API_KEY= -VOLCENGINE_CV_BASE_URL= +VOLCENGINE_CV_BASE_URL=https://ark.cn-beijing.volces.com/api/v3 + +# 火山方舟接入点 ID(在控制台「模型推理」创建接入点获取,形如 ep-2024xxxxxx-xxxxx) +# 图像模型(doubao-seedream)和视频模型(doubao-seedance)各创建一个接入点 +VOLCENGINE_IMAGE_ENDPOINT= +VOLCENGINE_VIDEO_ENDPOINT= diff --git a/examples/test_agnes.py b/examples/test_agnes.py new file mode 100644 index 0000000..7e5b6cf --- /dev/null +++ b/examples/test_agnes.py @@ -0,0 +1,74 @@ +"""临时:只测 agnes。测完可删。自动选模型。""" +import asyncio +from aibridge import Client + + +def load_env(): + env = {} + with open(".env") as f: + for line in f: + line = line.strip() + if line and not line.startswith("#") and "=" in line: + k, v = line.split("=", 1) + env[k.strip()] = v.strip() + return env + + +async def main(): + env = load_env() + c = Client( + provider="agnes", + api_key=env["AGNES_API_KEY"], + base_url=env["AGNES_BASE_URL"], + ) + await c.start() + + # 1. list_models + print("=== list_models ===") + models = [] + try: + models = await c.list_models() + print(f"✅ 通:{len(models)} 个模型") + for m in models: + print(f" - {m.id} ({m.type})") + except Exception as e: + print(f"❌ 失败: {e}") + + # 2. chat:自动选第一个 chat model + chat_m = next((m for m in models if str(m.type).lower() == "chat"), None) + if chat_m: + print(f"\n=== chat({chat_m.id}) ===") + try: + r = await c.chat( + model=chat_m.id, messages=[{"role": "user", "content": "说一个字"}] + ) + content = r.choices[0].message.content if r.choices else "(空)" + print(f"✅ 通: {str(content)[:30]}") + except Exception as e: + print(f"❌ 失败: {e}") + + # 3. image + image_m = next((m for m in models if str(m.type).lower() == "image"), None) + if image_m: + print(f"\n=== image({image_m.id}) ===") + try: + img = await c.image_generate(model=image_m.id, prompt="一只猫") + print(f"✅ 通: {'有数据' if img.data else '无数据'}") + except Exception as e: + print(f"❌ 失败: {e}") + + # 4. video + video_m = next((m for m in models if str(m.type).lower() == "video"), None) + if video_m: + print(f"\n=== video({video_m.id}) ===") + try: + task = await c.video_create(model=video_m.id, prompt="一只猫在走路") + print(f"✅ 通: task_id={task.task_id}") + except Exception as e: + print(f"❌ 失败: {e}") + + await c.close() + + +if __name__ == "__main__": + asyncio.run(main()) diff --git a/examples/test_volcengine.py b/examples/test_volcengine.py new file mode 100644 index 0000000..93a90a5 --- /dev/null +++ b/examples/test_volcengine.py @@ -0,0 +1,73 @@ +"""临时:只测火山引擎(国内可直连)。测完可删。 + +完全自动:list_models 后选第一个 image/video 模型测试,不硬编码 model 名。 +只报告每项通/不通。 +""" +import asyncio +from aibridge import Client + + +def load_env(): + env = {} + with open(".env") as f: + for line in f: + line = line.strip() + if line and not line.startswith("#") and "=" in line: + k, v = line.split("=", 1) + env[k.strip()] = v.strip() + return env + + +async def main(): + env = load_env() + c = Client( + provider="volcengine_cv", + api_key=env["VOLCENGINE_CV_API_KEY"], + base_url=env["VOLCENGINE_CV_BASE_URL"], + ) + await c.start() + + # 1. list_models(验证连通性 + 认证 + 代码) + print("=== list_models ===") + models = [] + try: + models = await c.list_models() + print(f"✅ 通:{len(models)} 个模型") + for t in ["chat", "image", "video", "embed"]: + m = next((m for m in models if str(m.type).lower() == t), None) + if m: + print(f" {t}: {m.id}") + except Exception as e: + print(f"❌ 失败: {e}") + await c.close() + return + + # 2. image:自动选第一个 image model + image_m = next((m for m in models if str(m.type).lower() == "image"), None) + if image_m: + print(f"\n=== image_generate({image_m.id}) ===") + try: + img = await c.image_generate(model=image_m.id, prompt="一只猫") + print(f"✅ 通:{'有数据' if img.data else '无数据'}") + except Exception as e: + print(f"❌ 失败: {e}") + else: + print("\nimage:list 无 image 模型,跳过") + + # 3. video:自动选第一个 video model + video_m = next((m for m in models if str(m.type).lower() == "video"), None) + if video_m: + print(f"\n=== video_create({video_m.id}) ===") + try: + task = await c.video_create(model=video_m.id, prompt="一只猫在走路") + print(f"✅ 通:task_id={task.task_id}") + except Exception as e: + print(f"❌ 失败: {e}") + else: + print("\nvideo:list 无 video 模型,跳过") + + await c.close() + + +if __name__ == "__main__": + asyncio.run(main()) From 6a0fa33d17eac9633a82ba8d56a45e9bb5c6323e Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Sat, 11 Jul 2026 07:45:43 +0800 Subject: [PATCH 54/55] =?UTF-8?q?fix:=20=E9=98=B6=E6=AE=B53=20CI=20?= =?UTF-8?q?=E4=BF=AE=E5=A4=8D=EF=BC=88.NET/Go/Python=20wheel=20=E6=9E=84?= =?UTF-8?q?=E5=BB=BA=EF=BC=89?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - .NET:DllImportSearchPath.Default→null、SetHandle→构造函数、Assembly using、DllImportResolver 签名 - Go:CGO LDFLAGS 路径修正(target/debug→target/release) - Python wheel:CI maturin build 路径修正(workspace target/wheels/) --- .github/workflows/ci.yml | 9 ++------- bindings/dotnet/AIBridge/Client.cs | 12 ++++-------- bindings/dotnet/AIBridge/Native.cs | 17 ++++++++++++----- bindings/go/aibridge.go | 2 +- 4 files changed, 19 insertions(+), 21 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index 356c2b3..f1be991 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -89,17 +89,12 @@ jobs: - name: maturin build wheel shell: bash run: | - cd crates/aibridge-python - if [ "${{ matrix.label }}" = "macos-universal2" ]; then - maturin build --release --universal2 - else - maturin build --release - fi + maturin build --manifest-path crates/aibridge-python/Cargo.toml --release $([ "${{ matrix.label }}" = "macos-universal2" ] && echo "--universal2") - name: 上传 wheel 产物 uses: actions/upload-artifact@v4 with: name: python-wheel-${{ matrix.label }} - path: crates/aibridge-python/target/wheels/*.whl + path: target/wheels/*.whl if-no-files-found: warn # ────────────────────────────────────────────────────────────────────────── diff --git a/bindings/dotnet/AIBridge/Client.cs b/bindings/dotnet/AIBridge/Client.cs index 10aaf62..cdc79dc 100644 --- a/bindings/dotnet/AIBridge/Client.cs +++ b/bindings/dotnet/AIBridge/Client.cs @@ -102,8 +102,7 @@ public ChatCompletion Chat(ChatRequest request) } // 成功:用 SafeHandle 接管字符串,拷贝后释放 - var handle = new AibridgeStringHandle(); - handle.SetHandle(outResponse); + var handle = new AibridgeStringHandle(outResponse); string? responseJson = handle.MarshalAndFree(); if (string.IsNullOrEmpty(responseJson)) @@ -169,8 +168,7 @@ public async IAsyncEnumerable ChatStreamAsync( if (s == AibridgeStatus.StreamChunk && outChunk != IntPtr.Zero) { // 拷贝 chunk JSON 并释放原生字符串(SafeHandle 兜底释放) - var h = new AibridgeStringHandle(); - h.SetHandle(outChunk); + var h = new AibridgeStringHandle(outChunk); json = h.MarshalAndFree(); } else if (outChunk != IntPtr.Zero) @@ -243,8 +241,7 @@ public SpeechResult Speech(SpeechRequest request) byte[] audioData = Array.Empty(); if (outAudio != IntPtr.Zero) { - var audioHandle = new AibridgeBytesHandle(); - audioHandle.SetHandle(outAudio); + var audioHandle = new AibridgeBytesHandle(outAudio); audioData = audioHandle.MarshalAndFree(); } @@ -252,8 +249,7 @@ public SpeechResult Speech(SpeechRequest request) SpeechResult? result; if (outMeta != IntPtr.Zero) { - var metaHandle = new AibridgeStringHandle(); - metaHandle.SetHandle(outMeta); + var metaHandle = new AibridgeStringHandle(outMeta); string? metaJson = metaHandle.MarshalAndFree(); result = string.IsNullOrEmpty(metaJson) ? new SpeechResult() diff --git a/bindings/dotnet/AIBridge/Native.cs b/bindings/dotnet/AIBridge/Native.cs index 9e0f87e..17959cd 100644 --- a/bindings/dotnet/AIBridge/Native.cs +++ b/bindings/dotnet/AIBridge/Native.cs @@ -1,4 +1,5 @@ using System.Runtime.InteropServices; +using System.Reflection; namespace AIBridge; @@ -61,6 +62,9 @@ internal sealed class AibridgeStringHandle : SafeHandle { public AibridgeStringHandle() : base(IntPtr.Zero, ownsHandle: true) { } + /// 用已有句柄构造(由调用方移交所有权)。 + public AibridgeStringHandle(IntPtr preexistingHandle) : base(preexistingHandle, ownsHandle: true) { } + public override bool IsInvalid => handle == IntPtr.Zero; protected override bool ReleaseHandle() @@ -89,6 +93,9 @@ internal sealed class AibridgeBytesHandle : SafeHandle { public AibridgeBytesHandle() : base(IntPtr.Zero, ownsHandle: true) { } + /// 用已有句柄构造(由调用方移交所有权)。 + public AibridgeBytesHandle(IntPtr preexistingHandle) : base(preexistingHandle, ownsHandle: true) { } + public override bool IsInvalid => handle == IntPtr.Zero; protected override bool ReleaseHandle() @@ -236,8 +243,8 @@ public static void Register(string libraryName) { if (Interlocked.CompareExchange(ref _registered, 1, 0) != 0) return; - // 解析回调签名:(string libName, Assembly asm, DllImportSearchPath? searchPath, IntPtr) => IntPtr - IntPtr Resolver(string lib, Assembly asm, DllImportSearchPath? search, IntPtr callers) + // 解析回调签名:(string libName, Assembly asm, DllImportSearchPath? searchPath) => IntPtr + IntPtr Resolver(string lib, Assembly asm, DllImportSearchPath? search) { // 仅处理本绑定的库名(其它库走默认解析) if (!string.Equals(lib, libraryName, StringComparison.OrdinalIgnoreCase)) @@ -267,7 +274,7 @@ IntPtr Resolver(string lib, Assembly asm, DllImportSearchPath? search, IntPtr ca string full = Path.Combine(dir, NativeFileName(libraryName)); if (File.Exists(full)) { - return NativeLibrary.Load(full, asm, DllImportSearchPath.Default); + return NativeLibrary.Load(full, asm, null); } } @@ -280,13 +287,13 @@ IntPtr Resolver(string lib, Assembly asm, DllImportSearchPath? search, IntPtr ca string full = Path.Combine(repoRoot, "target", profile, NativeFileName(libraryName)); if (File.Exists(full)) { - return NativeLibrary.Load(full, asm, DllImportSearchPath.Default); + return NativeLibrary.Load(full, asm, null); } } } // 最后让 .NET 用默认搜索(系统库目录等) - return NativeLibrary.Load(libraryName, asm, DllImportSearchPath.Default); + return NativeLibrary.Load(libraryName, asm, null); } NativeLibrary.SetDllImportResolver(typeof(Native).Assembly, Resolver); diff --git a/bindings/go/aibridge.go b/bindings/go/aibridge.go index c72bf86..9623bdf 100644 --- a/bindings/go/aibridge.go +++ b/bindings/go/aibridge.go @@ -16,7 +16,7 @@ package aibridge /* #cgo CFLAGS: -I${SRCDIR}/../../crates/aibridge-ffi/include -#cgo LDFLAGS: -L${SRCDIR}/../../target/debug -laibridge -lm +#cgo LDFLAGS: -L${SRCDIR}/../../target/release -laibridge -lm #include #include "aibridge.h" From 814b04269b5b85b9abd1d6b800c8060d6abf08b4 Mon Sep 17 00:00:00 2001 From: WingkySky <37488098@qq.com> Date: Mon, 13 Jul 2026 07:24:15 +0800 Subject: [PATCH 55/55] =?UTF-8?q?fix(ci):=20Go=20windows=20GNU=20=E5=B7=A5?= =?UTF-8?q?=E5=85=B7=E9=93=BE=20+=20Python=20universal2=20+=20=E7=A7=BB?= =?UTF-8?q?=E9=99=A4=20macos-x64?= MIME-Version: 1.0 Content-Type: text/plain; charset=UTF-8 Content-Transfer-Encoding: 8bit - Go windows-x64:Rust 改用 x86_64-pc-windows-gnu target(与 MinGW gcc 兼容,修复 __chkstk 链接失败) - Python macos-universal2:maturin --target universal2-apple-darwin(修复 --universal2 参数不存在) - 移除已弃用的 macos-x64 job(Rust 核心+ffi / Node .node,24h 超时) --- .github/workflows/ci.yml | 29 ++++++++++++++++++++++------- 1 file changed, 22 insertions(+), 7 deletions(-) diff --git a/.github/workflows/ci.yml b/.github/workflows/ci.yml index f1be991..784482a 100644 --- a/.github/workflows/ci.yml +++ b/.github/workflows/ci.yml @@ -7,7 +7,7 @@ # 说明: # - 不实际发布(无 token),仅验证跨平台构建 + 产出 artifact 供下载。 # - cargo 用 Swatinem/rust-cache 缓存;npm 由 setup-node 缓存;gradle 由 setup-gradle 缓存。 -# - ubuntu-24.04-arm 提供 linux arm64 runner;macos-13 提供 macOS amd64 (Intel) runner。 +# - ubuntu-24.04-arm 提供 linux arm64 runner;macos-x64 runner 已弃用(24h 超时),仅保留 macos-arm64。 name: CI @@ -37,8 +37,6 @@ jobs: label: linux-arm64 - os: macos-latest label: macos-arm64 - - os: macos-13 - label: macos-x64 - os: windows-latest label: windows-x64 steps: @@ -89,7 +87,13 @@ jobs: - name: maturin build wheel shell: bash run: | - maturin build --manifest-path crates/aibridge-python/Cargo.toml --release $([ "${{ matrix.label }}" = "macos-universal2" ] && echo "--universal2") + if [ "${{ matrix.label }}" = "macos-universal2" ]; then + # universal2: maturin 分别编译 x86_64/aarch64-apple-darwin 后 lipo 合并 + # (maturin 1.x 不支持 --universal2 标志,改用 --target universal2-apple-darwin) + maturin build --manifest-path crates/aibridge-python/Cargo.toml --release --target universal2-apple-darwin + else + maturin build --manifest-path crates/aibridge-python/Cargo.toml --release + fi - name: 上传 wheel 产物 uses: actions/upload-artifact@v4 with: @@ -113,8 +117,6 @@ jobs: label: linux-arm64 - os: macos-latest label: macos-arm64 - - os: macos-13 - label: macos-x64 - os: windows-latest label: windows-x64 steps: @@ -177,7 +179,20 @@ jobs: go-version: '1.26' cache: false - name: 构建 libaibridge (ffi, release) - run: cargo build -p aibridge-ffi --release + shell: bash + run: | + if [ "${{ matrix.label }}" = "windows-x64" ]; then + # Windows: 用 GNU target 编译,避免 MSVC 产 .lib 与 Go MinGW gcc 链接不兼容 + # (Go CGO 在 Windows 用 MinGW gcc;MSVC 运行时符号如 __chkstk 它找不到) + rustup target add x86_64-pc-windows-gnu + cargo build -p aibridge-ffi --release --target x86_64-pc-windows-gnu + # GNU target 产 libaibridge.dll.a(import lib)+ aibridge.dll,拷到 Go cgo 默认查找的 target/release + cp target/x86_64-pc-windows-gnu/release/libaibridge.dll.a target/release/ 2>/dev/null || \ + cp target/x86_64-pc-windows-gnu/release/libaibridge.a target/release/ 2>/dev/null || true + cp target/x86_64-pc-windows-gnu/release/aibridge.dll target/release/ 2>/dev/null || true + else + cargo build -p aibridge-ffi --release + fi - name: go build + vet (CGO) shell: bash env: