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:
leokun
2026-09-01 17:43:36 +08:00
parent 2c63bd845a
commit e768980dad
7 changed files with 396 additions and 32 deletions
@@ -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({
-1
View File
@@ -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)]
+7 -3
View File
@@ -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
View File
@@ -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, &params).await?;
let (request, cancellation, recorder) = self.request(request_id, &params).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, &params).await?;
let (request, cancellation, recorder) = self.request(request_id, &params).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", &params).await.unwrap();
let (_, _, second_recorder) = host.request("invocation", &params).await.unwrap();
let (_, body) = recorded_network_request(&params).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(&params).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", &params).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);
}
}
+67 -1
View File
@@ -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"
})]
);
}
}
+4 -2
View File
@@ -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())
}
}