Files
DevFlow/crates/df-ai/src/openai_helpers.rs
T

290 lines
12 KiB
Rust
Raw Blame History

This file contains ambiguous Unicode characters
This file contains Unicode characters that might be confused with other characters. If you think that this is intentional, you can safely ignore this warning. Use the Escape button to reveal them.
//! OpenAI 兼容协议适配 — 协议数据结构与 SSE 事件解析(纯函数)。
//!
//! 本模块从 `openai_compat.rs` 抽离,承载与 HTTP 无关的纯协议逻辑:
//! - 请求/响应结构体(`OpenAiRequest` / `OpenAiResponse` / `OpenAiStreamChunk` 等)
//! - SSE 事件 data → `StreamChunk` 转换纯函数(`apply_openai_sse`
//!
//! Provider struct + impl(含 HTTP 调用 / `convert_request` / `chat_url`)仍留在
//! `openai_compat.rs`Rust impl 块不可跨文件,故仅搬迁 impl 块外部的类型/纯函数。
//! 零行为变更(纯搬迁)。结构对齐 `anthropic_helpers.rs`。
use serde::{Deserialize, Serialize};
use tracing::{debug, error};
use crate::provider::{tool_call_id_or_fallback, StreamChunk, TokenUsage, ToolCallDelta};
// ============================================================
// OpenAI API 请求/响应结构体
// ============================================================
/// OpenAI 兼容请求体
#[derive(Debug, Clone, Serialize)]
pub(crate) struct OpenAiRequest {
pub model: String,
pub messages: Vec<OpenAiMessage>,
#[serde(skip_serializing_if = "Option::is_none")]
pub temperature: Option<f32>,
#[serde(skip_serializing_if = "Option::is_none")]
pub max_tokens: Option<u32>,
pub stream: bool,
#[serde(skip_serializing_if = "Option::is_none")]
pub tools: Option<Vec<serde_json::Value>>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_choice: Option<serde_json::Value>,
/// 流式时请求末 chunk 携带 usageOpenAI 官方 + DeepSeek/GLM 兼容)
#[serde(skip_serializing_if = "Option::is_none")]
pub stream_options: Option<serde_json::Value>,
/// DeepSeek thinking 模式
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
/// OpenAI 消息格式
///
/// `content` 为 `serde_json::Value`:纯文本消息走字符串简写(与老端点兼容),
/// 含图消息走 `[{type:"text",text},{type:"image_url",image_url:{url}}]` 数组。
#[derive(Debug, Clone, Serialize, Deserialize)]
pub(crate) struct OpenAiMessage {
pub role: String,
pub content: serde_json::Value,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_call_id: Option<String>,
#[serde(skip_serializing_if = "Option::is_none")]
pub tool_calls: Option<Vec<serde_json::Value>>,
/// DeepSeek thinking 模式推理内容
#[serde(skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
/// OpenAI 同步响应
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiResponse {
pub choices: Vec<OpenAiChoice>,
pub model: String,
pub usage: Option<OpenAiUsage>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiChoice {
pub message: OpenAiMessageResp,
// SW-260618-24: OpenAI 响应反序列化字段,保留以备调试/未来消费(如日志记录调用终止原因),标注意图消除 dead_code warning
#[allow(dead_code)]
pub finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiMessageResp {
pub content: Option<String>,
pub tool_calls: Option<Vec<OpenAiToolCallResp>>,
/// DeepSeek thinking 模式响应中的推理内容
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiToolCallResp {
pub id: String,
// SW-260618-24: 反序列化 #[serde(rename="type")] 字段,保留以对齐 OpenAI 响应结构,标注意图消除 dead_code warning
#[allow(dead_code)]
#[serde(rename = "type")]
pub call_type: String,
pub function: OpenAiFunctionResp,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiFunctionResp {
pub name: String,
pub arguments: String,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiUsage {
pub prompt_tokens: u32,
pub completion_tokens: u32,
pub total_tokens: u32,
/// DeepSeek 扩展:缓存命中 token(低价,deepseek-chat/reasoner prompt_cache_hit_tokens)。
/// OpenAI 官方(o1 等)无此字段 → serde default 0。其他 OpenAI 兼容网关若支持 cache 也用此名。
#[serde(default)]
pub prompt_cache_hit_tokens: u32,
/// DeepSeek 扩展:未命中 token(全价真实输入,prompt_cache_miss_tokens)。
/// OpenAI 官方无此字段 → serde default 0。
#[serde(default)]
pub prompt_cache_miss_tokens: u32,
/// DeepSeek-reasoner / OpenAI o1 扩展:思考 token(隐藏输出,reasoning_tokens)。
/// 非 reasoning 模型无此字段 → serde default 0。
#[serde(default)]
pub reasoning_tokens: u32,
}
/// SSE 流式响应 chunk
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiStreamChunk {
pub choices: Vec<OpenAiStreamChoice>,
/// 末 chunkchoices 为空)携带的累计 usage
#[serde(default)]
pub usage: Option<OpenAiUsage>,
/// 流中途 error 事件(OpenAI 兼容协议:`{"error":{"message":..,"type":..}}`)。
/// 部分中转站按 OpenAI 协议在流中途发 error 帧而非走 HTTP 非 200
/// serde default + Value 兜底:旧响应无此字段不受影响,且对 error 载荷形态不敏感。
#[serde(default)]
pub error: Option<serde_json::Value>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiStreamChoice {
pub delta: OpenAiStreamDelta,
pub finish_reason: Option<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiStreamDelta {
pub content: Option<String>,
pub tool_calls: Option<Vec<OpenAiStreamToolCall>>,
/// DeepSeek thinking 模式流式推理内容
#[serde(default, skip_serializing_if = "Option::is_none")]
pub reasoning_content: Option<String>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiStreamToolCall {
pub index: u32,
pub id: Option<String>,
pub function: Option<OpenAiStreamFunction>,
}
#[derive(Debug, Deserialize)]
pub(crate) struct OpenAiStreamFunction {
pub name: Option<String>,
pub arguments: Option<String>,
}
// ============================================================
// SSE 解析纯函数(与 HTTP 解耦,便于单测)
// ============================================================
/// 将一条 OpenAI 兼容 SSE 事件 data 解析为 StreamChunk,并按需更新 usage 累加器。
///
/// - `[DONE]` → 返回 `finished=true` 的终态 chunk`usage` 取自累加器(`take()`)。
/// - 普通文本/工具增量 chunk → 返回对应 `StreamChunk`usage 字段恒为 Noneusage 仅在终态带出)。
/// - usage`stream_options.include_usage` 时末段或 usage-only chunk 携带)→ 覆盖累加器(覆盖语义保对)。
/// - 解析失败 → 返回空 chunk(与原内联实现一致)。
///
/// 等价性:delta / tool_calls / finished / 解析失败等分支与原 stream() 闭包逐字一致;
/// usage 透传([DONE] 终态 take() 带出、usage chunk 覆盖累加器)为本次新增能力,
/// 对应 StreamChunk 新增的 usage 字段 + 请求体新增 stream_options.include_usage。
pub(crate) fn apply_openai_sse(data: &str, usage_accum: &mut Option<TokenUsage>) -> StreamChunk {
// OpenAI 发送 "data: [DONE]" 表示流结束,带出累积 usage
if data == "[DONE]" {
return StreamChunk {
delta: String::new(),
finished: true,
tool_calls: None,
usage: usage_accum.take(),
error: None,
reasoning_content: None,
};
}
match serde_json::from_str::<OpenAiStreamChunk>(data) {
Ok(chunk) => {
// 流中途 error 事件(中转站按 OpenAI 协议在流中途发 error 帧)。
// 不走 finished 完成路径(避免残缺响应被当正常完成入库),由 stream_llm
// 识别 error 非空 → 发 AiError + 丢弃残缺(对齐 anthropic_helpers 215-219)。
if let Some(err_val) = chunk.error {
let msg = err_val
.get("message")
.and_then(|m| m.as_str())
.unwrap_or("stream error")
.to_string();
error!(%msg, raw = %err_val, "OpenAI 流式错误事件");
return StreamChunk {
delta: String::new(),
finished: false,
tool_calls: None,
usage: None,
error: Some(msg),
reasoning_content: None,
};
}
// 提取 usage(带 include_usage 时末段 chunk 携带,覆盖累积)
if let Some(u) = chunk.usage {
tracing::info!(
prompt = u.prompt_tokens,
completion = u.completion_tokens,
cache_hit = u.prompt_cache_hit_tokens,
cache_miss = u.prompt_cache_miss_tokens,
reasoning = u.reasoning_tokens,
"[OpenAI] 末 chunk usage 解析(deepseek 等报 cache)"
);
*usage_accum = Some(TokenUsage {
prompt_tokens: u.prompt_tokens,
completion_tokens: u.completion_tokens,
total_tokens: u.total_tokens,
prompt_cache_hit_tokens: u.prompt_cache_hit_tokens,
prompt_cache_miss_tokens: u.prompt_cache_miss_tokens,
reasoning_tokens: u.reasoning_tokens,
});
}
if let Some(choice) = chunk.choices.into_iter().next() {
let delta_text = choice.delta.content.unwrap_or_default();
// "length" = max_tokens 截断,属正常终止(非断连),纳入 finished
let finished = choice.finish_reason.as_deref() == Some("stop")
|| choice.finish_reason.as_deref() == Some("tool_calls")
|| choice.finish_reason.as_deref() == Some("length");
let tool_calls = choice.delta.tool_calls.map(|tcs| {
tcs.into_iter()
.map(|tc| {
// CR-空 id:流式 chunk 的 id 可能为 Some("")SenseNova 兼容缺陷)。
// 仅对「provider 显式给了 id 字段」的 chunk 做兜底——NoneOpenAI
// 协议:仅首 chunk 携带 id,后续 chunk 无 id)保持 None,避免
// 覆盖首 chunk 的权威 id。Some("") → `gen_stream_{index}` fallback
// Some(非空) → 原样。下游 stream_recv 按 index 累积,draft.id 透传
// 至 ToolCall.idaccumulate_tool_calls 仅 Some 覆盖,None 不动)。
let id = tc.id.map(|raw| tool_call_id_or_fallback(&raw, tc.index as usize, "gen_stream"));
ToolCallDelta {
index: tc.index,
id,
function_name: tc.function.as_ref().and_then(|f| f.name.clone()),
function_arguments: tc.function.and_then(|f| f.arguments),
}
})
.collect()
});
StreamChunk {
delta: delta_text,
finished,
tool_calls,
// 加固:usage 与 finish_reason 同帧的端点(部分兼容实现),此处已累积则挂上,
// 使下游 stream_recv 不必等到 [DONE] 帧即可拿到真实 usage。
usage: usage_accum.clone(),
error: None,
reasoning_content: choice.delta.reasoning_content,
}
} else {
// choices 为空 = usage-only chunk,不输出文本(usage 已累积),usage 一并带出
StreamChunk {
delta: String::new(),
finished: false,
tool_calls: None,
usage: usage_accum.clone(),
error: None,
reasoning_content: None,
}
}
}
Err(e) => {
debug!("SSE 数据解析失败: {} — data: {}", e, data);
StreamChunk {
delta: String::new(),
finished: false,
tool_calls: None,
usage: None,
error: None,
reasoning_content: None,
}
}
}
}