From e768980dadd96c8e96a6af483d53be6e82c2f591 Mon Sep 17 00:00:00 2001 From: leokun Date: Tue, 1 Sep 2026 17:43:36 +0800 Subject: [PATCH] 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. --- .../plugins/build-in/codex-auth/codex_test.ts | 38 +++ server/src/plugin/mod.rs | 1 - server/src/plugin/registry.rs | 10 +- .../plugin/sdk/protocol/openai_responses.ts | 12 +- server/src/plugin/worker.rs | 293 ++++++++++++++++-- server/src/provider/openai_responses.rs | 68 +++- server/src/provider/router.rs | 6 +- 7 files changed, 396 insertions(+), 32 deletions(-) diff --git a/server/plugins/build-in/codex-auth/codex_test.ts b/server/plugins/build-in/codex-auth/codex_test.ts index b572cdd..1149537 100644 --- a/server/plugins/build-in/codex-auth/codex_test.ts +++ b/server/plugins/build-in/codex-auth/codex_test.ts @@ -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({ diff --git a/server/src/plugin/mod.rs b/server/src/plugin/mod.rs index ce7faa3..7789c40 100644 --- a/server/src/plugin/mod.rs +++ b/server/src/plugin/mod.rs @@ -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)] diff --git a/server/src/plugin/registry.rs b/server/src/plugin/registry.rs index e107caf..29e22bc 100644 --- a/server/src/plugin/registry.rs +++ b/server/src/plugin/registry.rs @@ -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 { diff --git a/server/src/plugin/sdk/protocol/openai_responses.ts b/server/src/plugin/sdk/protocol/openai_responses.ts index b8e479a..b04e037 100644 --- a/server/src/plugin/sdk/protocol/openai_responses.ts +++ b/server/src/plugin/sdk/protocol/openai_responses.ts @@ -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 = { 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 { diff --git a/server/src/plugin/worker.rs b/server/src/plugin/worker.rs index 2863b04..b56c867 100644 --- a/server/src/plugin/worker.rs +++ b/server/src/plugin/worker.rs @@ -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>, } +struct InvocationState { + cancellation: CancellationToken, + recorder: Option, + recorder_claimed: AtomicBool, +} + +impl InvocationState { + fn claim_recorder(&self) -> Option { + 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>, store: Store, - cancellations: Arc>>, + invocations: Arc>>>, streams: Arc>>, } @@ -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 { - 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, ) -> Result> { 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>> { @@ -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, + )> { 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 { - 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 { - 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::>(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::>(); @@ -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::(), + 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); + } +} diff --git a/server/src/provider/openai_responses.rs b/server/src/provider/openai_responses.rs index 99eb7a7..6439ea4 100644 --- a/server/src/provider/openai_responses.rs +++ b/server/src/provider/openai_responses.rs @@ -420,7 +420,12 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result> { .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::>>()?, + ); } push_responses_text(&mut input, &message.role, text); for call in calls { @@ -437,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result> { Ok(input) } +fn response_reasoning_input(item: &Value) -> Result { + 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, 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" + })] + ); + } +} diff --git a/server/src/provider/router.rs b/server/src/provider/router.rs index d16f122..7eabe64 100644 --- a/server/src/provider/router.rs +++ b/server/src/provider/router.rs @@ -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 = 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()) } }