mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:40:50 +08:00
feat: enhance reasoning replay functionality and integrate call recording
- Added a new test to validate the projection of reasoning response items to valid input items in the Codex API. - Introduced `CallRecorder` to track network requests and responses during plugin interactions. - Updated the `PluginRegistry` and `PluginWorker` to support call recording, ensuring that reasoning items are correctly processed and recorded. - Refactored the `responses_input` function to handle reasoning items more effectively, improving the overall response handling logic.
This commit is contained in:
@@ -8,6 +8,7 @@ import type { LlmRequest, ModelEvent } from "cursor-byok:provider";
|
||||
import type { ResourceSnapshot } from "cursor-byok:resource";
|
||||
import { codexDeviceOAuth } from "./oauth.ts";
|
||||
import { parseOfficialModels } from "./models.ts";
|
||||
import { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
|
||||
import { codexProvider, isQuotaError } from "./provider.ts";
|
||||
import {
|
||||
accountIdentity,
|
||||
@@ -305,6 +306,43 @@ Deno.test("invoke streams normalized events from the Codex Responses API", async
|
||||
]);
|
||||
});
|
||||
|
||||
Deno.test("reasoning replay projects response items to valid input items", () => {
|
||||
const replayRequest = request();
|
||||
replayRequest.messages = [{
|
||||
role: "assistant",
|
||||
text: "",
|
||||
thinking: "",
|
||||
replayState: {
|
||||
providerKind: "openai_responses",
|
||||
value: {
|
||||
items: [{
|
||||
type: "reasoning",
|
||||
id: "item-1",
|
||||
status: "completed",
|
||||
summary: [{ type: "summary_text", text: "why" }],
|
||||
content: [],
|
||||
encrypted_content: "opaque",
|
||||
output_only: true,
|
||||
}],
|
||||
},
|
||||
},
|
||||
toolCalls: [],
|
||||
}];
|
||||
|
||||
const body = buildResponsesBody({
|
||||
url: "https://example.com/responses",
|
||||
model: "gpt-test",
|
||||
request: replayRequest,
|
||||
});
|
||||
assertEquals(body.input, [{
|
||||
type: "reasoning",
|
||||
id: "item-1",
|
||||
summary: [{ type: "summary_text", text: "why" }],
|
||||
content: [],
|
||||
encrypted_content: "opaque",
|
||||
}]);
|
||||
});
|
||||
|
||||
Deno.test("invoke streams incremental tool calls and replays reasoning items", async () => {
|
||||
const token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||
const draft = await credentialDraft({
|
||||
|
||||
@@ -20,7 +20,6 @@ pub use descriptor::{
|
||||
};
|
||||
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
||||
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
||||
pub(crate) use wire::llm_request as plugin_llm_request;
|
||||
|
||||
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
||||
#[cfg(windows)]
|
||||
|
||||
@@ -20,8 +20,11 @@ use super::{
|
||||
worker::{PluginWorker, WorkerStreamItem},
|
||||
};
|
||||
use crate::{
|
||||
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
|
||||
Result,
|
||||
model::ModelInvocation,
|
||||
provider::ProviderStream,
|
||||
provider::{CallRecorder, ModelEvent},
|
||||
store::Store,
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
||||
@@ -203,6 +206,7 @@ impl PluginRegistry {
|
||||
&self,
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
recorder: CallRecorder,
|
||||
) -> ProviderStream {
|
||||
let registry = self.clone();
|
||||
Box::pin(try_stream! {
|
||||
@@ -232,7 +236,7 @@ impl PluginRegistry {
|
||||
"request": request,
|
||||
});
|
||||
let worker = registry.worker(&entry, &executable).await;
|
||||
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone()).await?;
|
||||
let mut items = worker.invoke_streaming("provider.invoke", params, cancellation.clone(), Some(recorder)).await?;
|
||||
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
|
||||
while let Some(item) = items.recv().await {
|
||||
match item {
|
||||
|
||||
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
|
||||
if (!Array.isArray(items)) {
|
||||
throw new Error("OpenAI Responses replay state is missing items");
|
||||
}
|
||||
return items;
|
||||
return items.map((item) => {
|
||||
const source = record(item);
|
||||
if (source?.type !== "reasoning") {
|
||||
throw new Error("OpenAI Responses replay state contains a non-reasoning item");
|
||||
}
|
||||
const projected: Record<string, JsonValue> = { type: "reasoning" };
|
||||
for (const field of ["id", "summary", "content", "encrypted_content"] as const) {
|
||||
if (field in source) projected[field] = source[field] as JsonValue;
|
||||
}
|
||||
return projected;
|
||||
});
|
||||
}
|
||||
|
||||
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
|
||||
|
||||
+269
-24
@@ -3,7 +3,10 @@ use std::{
|
||||
collections::{HashMap, HashSet},
|
||||
path::PathBuf,
|
||||
process::Stdio,
|
||||
sync::Arc,
|
||||
sync::{
|
||||
atomic::{AtomicBool, Ordering},
|
||||
Arc,
|
||||
},
|
||||
time::Duration,
|
||||
};
|
||||
|
||||
@@ -19,7 +22,7 @@ use super::{
|
||||
definition::{file_url, PluginDefinitionLoader},
|
||||
protocol::{HostMessage, WorkerMessage},
|
||||
};
|
||||
use crate::{store::Store, Error, Result};
|
||||
use crate::{provider::CallRecorder, store::Store, Error, Result};
|
||||
|
||||
const INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
||||
@@ -56,12 +59,29 @@ struct WorkerProcess {
|
||||
stdin: Arc<Mutex<ChildStdin>>,
|
||||
}
|
||||
|
||||
struct InvocationState {
|
||||
cancellation: CancellationToken,
|
||||
recorder: Option<CallRecorder>,
|
||||
recorder_claimed: AtomicBool,
|
||||
}
|
||||
|
||||
impl InvocationState {
|
||||
fn claim_recorder(&self) -> Option<CallRecorder> {
|
||||
self.recorder.as_ref().and_then(|recorder| {
|
||||
self.recorder_claimed
|
||||
.compare_exchange(false, true, Ordering::AcqRel, Ordering::Acquire)
|
||||
.ok()
|
||||
.map(|_| recorder.clone())
|
||||
})
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
struct HostContext {
|
||||
plugin_id: String,
|
||||
network_hosts: Arc<HashSet<String>>,
|
||||
store: Store,
|
||||
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
|
||||
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
|
||||
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
||||
}
|
||||
|
||||
@@ -87,7 +107,7 @@ impl PluginWorker {
|
||||
.collect(),
|
||||
),
|
||||
store,
|
||||
cancellations: Arc::new(Mutex::new(HashMap::new())),
|
||||
invocations: Arc::new(Mutex::new(HashMap::new())),
|
||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||
},
|
||||
plugin_id,
|
||||
@@ -108,7 +128,9 @@ impl PluginWorker {
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
) -> Result<serde_json::Value> {
|
||||
let mut items = self.invoke_streaming(method, params, cancellation).await?;
|
||||
let mut items = self
|
||||
.invoke_streaming(method, params, cancellation, None)
|
||||
.await?;
|
||||
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
|
||||
while let Some(item) = items.recv().await {
|
||||
if let WorkerStreamItem::Result(result) = item {
|
||||
@@ -137,15 +159,18 @@ impl PluginWorker {
|
||||
method: &str,
|
||||
params: serde_json::Value,
|
||||
cancellation: CancellationToken,
|
||||
recorder: Option<CallRecorder>,
|
||||
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
||||
let id = uuid::Uuid::new_v4().to_string();
|
||||
let request_cancellation = CancellationToken::new();
|
||||
self.inner
|
||||
.host
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.insert(id.clone(), request_cancellation.clone());
|
||||
self.inner.host.invocations.lock().await.insert(
|
||||
id.clone(),
|
||||
Arc::new(InvocationState {
|
||||
cancellation: request_cancellation.clone(),
|
||||
recorder,
|
||||
recorder_claimed: AtomicBool::new(false),
|
||||
}),
|
||||
);
|
||||
let (sender, receiver) = mpsc::unbounded_channel();
|
||||
self.inner
|
||||
.pending
|
||||
@@ -182,10 +207,10 @@ impl PluginWorker {
|
||||
}
|
||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
||||
inner.pending.lock().await.remove(&request_id);
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
inner.host.invocations.lock().await.remove(&request_id);
|
||||
}
|
||||
_ = sender.closed() => {
|
||||
inner.host.cancellations.lock().await.remove(&request_id);
|
||||
inner.host.invocations.lock().await.remove(&request_id);
|
||||
}
|
||||
}
|
||||
});
|
||||
@@ -201,7 +226,7 @@ impl PluginWorker {
|
||||
|
||||
async fn cleanup(&self, id: &str) {
|
||||
self.inner.pending.lock().await.remove(id);
|
||||
self.inner.host.cancellations.lock().await.remove(id);
|
||||
self.inner.host.invocations.lock().await.remove(id);
|
||||
}
|
||||
|
||||
async fn stdin(&self) -> Result<Arc<Mutex<ChildStdin>>> {
|
||||
@@ -393,6 +418,28 @@ async fn fail_pending(pending: &Pending, message: &str) {
|
||||
}
|
||||
}
|
||||
|
||||
fn recorded_network_request(
|
||||
params: &serde_json::Value,
|
||||
) -> Result<(serde_json::Value, serde_json::Value)> {
|
||||
let mut recorded_headers = serde_json::Map::new();
|
||||
if let Some(headers) = params.get("headers").and_then(serde_json::Value::as_object) {
|
||||
for (name, value) in headers {
|
||||
let value = value.as_str().ok_or_else(|| {
|
||||
Error::Config(format!("plugin HTTP header '{name}' must be a string"))
|
||||
})?;
|
||||
if !crate::model::is_sensitive_header(name) {
|
||||
recorded_headers.insert(name.clone(), value.into());
|
||||
}
|
||||
}
|
||||
}
|
||||
let body = params
|
||||
.get("body")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(|body| serde_json::from_str(body).unwrap_or_else(|_| body.into()))
|
||||
.unwrap_or(serde_json::Value::Null);
|
||||
Ok((serde_json::Value::Object(recorded_headers), body))
|
||||
}
|
||||
|
||||
impl HostContext {
|
||||
async fn call(
|
||||
&self,
|
||||
@@ -421,7 +468,11 @@ impl HostContext {
|
||||
&self,
|
||||
request_id: &str,
|
||||
params: &serde_json::Value,
|
||||
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
|
||||
) -> Result<(
|
||||
reqwest::RequestBuilder,
|
||||
CancellationToken,
|
||||
Option<CallRecorder>,
|
||||
)> {
|
||||
let raw_url = required_string(params, "url")?;
|
||||
let url = url::Url::parse(raw_url)
|
||||
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
|
||||
@@ -463,14 +514,17 @@ impl HostContext {
|
||||
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
|
||||
request = request.body(body.to_owned());
|
||||
}
|
||||
let cancellation = self
|
||||
.cancellations
|
||||
.lock()
|
||||
.await
|
||||
.get(request_id)
|
||||
.cloned()
|
||||
let invocation = self.invocations.lock().await.get(request_id).cloned();
|
||||
let cancellation = invocation
|
||||
.as_ref()
|
||||
.map(|state| state.cancellation.clone())
|
||||
.unwrap_or_default();
|
||||
Ok((request, cancellation))
|
||||
let recorder = invocation.and_then(|state| state.claim_recorder());
|
||||
if let Some(recorder) = &recorder {
|
||||
let (headers, body) = recorded_network_request(params)?;
|
||||
recorder.request(headers, &body).await?;
|
||||
}
|
||||
Ok((request, cancellation, recorder))
|
||||
}
|
||||
|
||||
async fn fetch(
|
||||
@@ -478,13 +532,16 @@ impl HostContext {
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let request = request.timeout(Duration::from_secs(60));
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_headers(status).await?;
|
||||
}
|
||||
if response
|
||||
.content_length()
|
||||
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
||||
@@ -503,6 +560,9 @@ impl HostContext {
|
||||
"plugin network response is larger than allowed".into(),
|
||||
));
|
||||
}
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_chunk(&body).await?;
|
||||
}
|
||||
Ok(
|
||||
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
||||
)
|
||||
@@ -514,12 +574,15 @@ impl HostContext {
|
||||
request_id: &str,
|
||||
params: serde_json::Value,
|
||||
) -> Result<serde_json::Value> {
|
||||
let (request, cancellation) = self.request(request_id, ¶ms).await?;
|
||||
let (request, cancellation, recorder) = self.request(request_id, ¶ms).await?;
|
||||
let response = tokio::select! {
|
||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||
response = request.send() => response?,
|
||||
};
|
||||
let status = response.status().as_u16();
|
||||
if let Some(recorder) = &recorder {
|
||||
recorder.response_headers(status).await?;
|
||||
}
|
||||
let headers = header_map(&response);
|
||||
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
||||
tokio::spawn(async move {
|
||||
@@ -552,6 +615,12 @@ impl HostContext {
|
||||
.await;
|
||||
return;
|
||||
}
|
||||
if let Some(recorder) = &recorder {
|
||||
if let Err(error) = recorder.response_chunk(&chunk).await {
|
||||
let _ = sender.send(Err(error)).await;
|
||||
return;
|
||||
}
|
||||
}
|
||||
buffered.extend_from_slice(&chunk);
|
||||
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
||||
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
|
||||
@@ -645,3 +714,179 @@ fn required_string<'a>(params: &'a serde_json::Value, key: &str) -> Result<&'a s
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use crate::{
|
||||
model::{NewLlmCall, ProviderType},
|
||||
provider::{CallRecorder, FinishReason},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::*;
|
||||
|
||||
async fn recorder(detailed: bool, call_id: &str) -> (tempfile::TempDir, Store, CallRecorder) {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("test.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
store.set_detailed_logging(detailed).await.unwrap();
|
||||
let recorder = CallRecorder::start(
|
||||
store.clone(),
|
||||
NewLlmCall {
|
||||
call_id: call_id.into(),
|
||||
run_id: "run".into(),
|
||||
conversation_id: "conversation".into(),
|
||||
provider_call_index: 0,
|
||||
model_hash: "plugin:test/provider/model".into(),
|
||||
provider_type: ProviderType::Plugin,
|
||||
provider_url: "plugin://test/provider".into(),
|
||||
request_type: ProviderType::Plugin,
|
||||
request_url: "plugin://test/provider".into(),
|
||||
model_id: "model".into(),
|
||||
display_name: "Model".into(),
|
||||
reasoning_effort: None,
|
||||
fast: false,
|
||||
message_count: 1,
|
||||
tool_count: 0,
|
||||
detailed: false,
|
||||
},
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
(directory, store, recorder)
|
||||
}
|
||||
|
||||
fn network_params() -> serde_json::Value {
|
||||
serde_json::json!({
|
||||
"url": "https://example.com/v1/responses",
|
||||
"method": "POST",
|
||||
"headers": {
|
||||
"Authorization": "Bearer secret",
|
||||
"X-Api-Key": "secret-key",
|
||||
"Cookie": "session=secret",
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
},
|
||||
"body": "{\"model\":\"test\",\"stream\":true}"
|
||||
})
|
||||
}
|
||||
|
||||
async fn host_with_recorder(store: Store, recorder: CallRecorder) -> HostContext {
|
||||
let invocations = Arc::new(Mutex::new(HashMap::new()));
|
||||
invocations.lock().await.insert(
|
||||
"invocation".into(),
|
||||
Arc::new(InvocationState {
|
||||
cancellation: CancellationToken::new(),
|
||||
recorder: Some(recorder),
|
||||
recorder_claimed: AtomicBool::new(false),
|
||||
}),
|
||||
);
|
||||
HostContext {
|
||||
plugin_id: "test".into(),
|
||||
network_hosts: Arc::new(HashSet::from(["example.com".into()])),
|
||||
store,
|
||||
invocations,
|
||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn recorded_plugin_request_omits_sensitive_headers() {
|
||||
let (headers, body) = recorded_network_request(&network_params()).unwrap();
|
||||
|
||||
assert_eq!(
|
||||
headers,
|
||||
serde_json::json!({
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
})
|
||||
);
|
||||
assert_eq!(body, serde_json::json!({ "model": "test", "stream": true }));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn detailed_plugin_network_recording_persists_request_and_raw_response() {
|
||||
let (_directory, store, recorder) = recorder(true, "detailed-plugin").await;
|
||||
let host = host_with_recorder(store.clone(), recorder.clone()).await;
|
||||
let params = network_params();
|
||||
let (_, _, first_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, _, second_recorder) = host.request("invocation", ¶ms).await.unwrap();
|
||||
let (_, body) = recorded_network_request(¶ms).unwrap();
|
||||
|
||||
assert!(first_recorder.is_some());
|
||||
assert!(second_recorder.is_none());
|
||||
recorder.response_headers(200).await.unwrap();
|
||||
recorder
|
||||
.response_chunk(b"data: {\"type\":\"response.created\"}\n\n")
|
||||
.await
|
||||
.unwrap();
|
||||
recorder.response_chunk(b"data: [DONE]\n\n").await.unwrap();
|
||||
recorder.completed(FinishReason::Stop).await.unwrap();
|
||||
|
||||
let request = store
|
||||
.llm_call_request("detailed-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.unwrap();
|
||||
assert_eq!(
|
||||
request.headers,
|
||||
serde_json::json!({
|
||||
"content-type": "application/json",
|
||||
"x-client-request-id": "request-1"
|
||||
})
|
||||
);
|
||||
assert_eq!(request.body, body);
|
||||
let chunks = store.llm_call_chunks("detailed-plugin").await.unwrap();
|
||||
let expected_response = "data: {\"type\":\"response.created\"}\n\ndata: [DONE]\n\n";
|
||||
assert_eq!(chunks.len(), 2);
|
||||
assert_eq!(
|
||||
chunks
|
||||
.iter()
|
||||
.map(|chunk| chunk.data.as_str())
|
||||
.collect::<String>(),
|
||||
expected_response
|
||||
);
|
||||
let summary = store.llm_call("detailed-plugin").await.unwrap().unwrap();
|
||||
assert_eq!(summary.http_status, Some(200));
|
||||
assert_eq!(summary.stream_event_count, 2);
|
||||
assert_eq!(summary.response_bytes, expected_response.len() as i64);
|
||||
assert!(summary.detailed);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn standard_plugin_network_recording_keeps_metrics_without_payloads() {
|
||||
let (_directory, store, recorder) = recorder(false, "standard-plugin").await;
|
||||
let host = host_with_recorder(store.clone(), recorder.clone()).await;
|
||||
let params = network_params();
|
||||
let (_, body) = recorded_network_request(¶ms).unwrap();
|
||||
let request_bytes = serde_json::to_string(&body).unwrap().len() as i64;
|
||||
let response = b"data: [DONE]\n\n";
|
||||
|
||||
let (_, _, observed) = host.request("invocation", ¶ms).await.unwrap();
|
||||
assert!(observed.is_some());
|
||||
recorder.response_headers(204).await.unwrap();
|
||||
recorder.response_chunk(response).await.unwrap();
|
||||
recorder.completed(FinishReason::Stop).await.unwrap();
|
||||
|
||||
assert!(store
|
||||
.llm_call_request("standard-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none());
|
||||
assert!(store
|
||||
.llm_call_chunks("standard-plugin")
|
||||
.await
|
||||
.unwrap()
|
||||
.is_empty());
|
||||
let summary = store.llm_call("standard-plugin").await.unwrap().unwrap();
|
||||
assert_eq!(summary.http_status, Some(204));
|
||||
assert_eq!(summary.request_bytes, Some(request_bytes));
|
||||
assert_eq!(summary.response_bytes, response.len() as i64);
|
||||
assert_eq!(summary.stream_event_count, 1);
|
||||
assert!(!summary.detailed);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -420,7 +420,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
|
||||
.ok_or_else(|| {
|
||||
Error::Protocol("OpenAI Responses replay state is missing items".into())
|
||||
})?;
|
||||
input.extend(items.iter().cloned());
|
||||
input.extend(
|
||||
items
|
||||
.iter()
|
||||
.map(response_reasoning_input)
|
||||
.collect::<Result<Vec<_>>>()?,
|
||||
);
|
||||
}
|
||||
push_responses_text(&mut input, &message.role, text);
|
||||
for call in calls {
|
||||
@@ -437,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
|
||||
Ok(input)
|
||||
}
|
||||
|
||||
fn response_reasoning_input(item: &Value) -> Result<Value> {
|
||||
let source = item
|
||||
.as_object()
|
||||
.filter(|object| object.get("type").and_then(Value::as_str) == Some("reasoning"))
|
||||
.ok_or_else(|| {
|
||||
Error::Protocol("OpenAI Responses replay state contains a non-reasoning item".into())
|
||||
})?;
|
||||
let mut projected = Map::new();
|
||||
projected.insert("type".into(), json!("reasoning"));
|
||||
for field in ["id", "summary", "content", "encrypted_content"] {
|
||||
if let Some(value) = source.get(field) {
|
||||
projected.insert(field.into(), value.clone());
|
||||
}
|
||||
}
|
||||
Ok(Value::Object(projected))
|
||||
}
|
||||
|
||||
fn push_responses_parts(input: &mut Vec<Value>, role: &Role, parts: &[ContentPart]) -> Result<()> {
|
||||
let text_type = if *role == Role::Assistant {
|
||||
"output_text"
|
||||
@@ -517,3 +539,47 @@ fn responses_usage(value: &Value) -> Usage {
|
||||
.and_then(Value::as_u64),
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
use crate::model::ProviderReplayState;
|
||||
|
||||
#[test]
|
||||
fn reasoning_replay_projects_response_items_to_valid_input_items() {
|
||||
let messages = [ProjectedMessage {
|
||||
message_id: "assistant-1".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: String::new(),
|
||||
replay_state: Some(ProviderReplayState {
|
||||
provider_kind: "openai_responses".into(),
|
||||
value: json!({
|
||||
"items": [{
|
||||
"type": "reasoning",
|
||||
"id": "item-1",
|
||||
"status": "completed",
|
||||
"summary": [{"type": "summary_text", "text": "why"}],
|
||||
"content": [],
|
||||
"encrypted_content": "opaque",
|
||||
"output_only": true
|
||||
}]
|
||||
}),
|
||||
}),
|
||||
calls: Vec::new(),
|
||||
},
|
||||
}];
|
||||
|
||||
assert_eq!(
|
||||
responses_input(&messages).unwrap(),
|
||||
vec![json!({
|
||||
"type": "reasoning",
|
||||
"id": "item-1",
|
||||
"summary": [{"type": "summary_text", "text": "why"}],
|
||||
"content": [],
|
||||
"encrypted_content": "opaque"
|
||||
})]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -62,7 +62,6 @@ impl Provider for ProviderRouter {
|
||||
let plan = plugins.plan_model(&selected).await?;
|
||||
let recorder = start_recorder(&store, &invocation, &selected, &plan.model.display_name, ProviderType::Plugin, &plan.request_url, &plan.model.model_id).await?;
|
||||
let guard = recorder.cancel_on_drop();
|
||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
||||
let mut routed = invocation.clone();
|
||||
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
||||
if let Some(tokens) = plan.model.max_output_tokens {
|
||||
@@ -70,6 +69,7 @@ impl Provider for ProviderRouter {
|
||||
}
|
||||
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||
registry: plugins.clone(),
|
||||
recorder: recorder.clone(),
|
||||
})));
|
||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||
} else {
|
||||
@@ -227,6 +227,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken
|
||||
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
||||
struct PluginModelProvider {
|
||||
registry: PluginRegistry,
|
||||
recorder: CallRecorder,
|
||||
}
|
||||
|
||||
impl Provider for PluginModelProvider {
|
||||
@@ -235,7 +236,8 @@ impl Provider for PluginModelProvider {
|
||||
invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
self.registry.stream_model(invocation, cancellation)
|
||||
self.registry
|
||||
.stream_model(invocation, cancellation, self.recorder.clone())
|
||||
}
|
||||
}
|
||||
|
||||
|
||||
Reference in New Issue
Block a user