mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-09 00:16:23 +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 type { ResourceSnapshot } from "cursor-byok:resource";
|
||||||
import { codexDeviceOAuth } from "./oauth.ts";
|
import { codexDeviceOAuth } from "./oauth.ts";
|
||||||
import { parseOfficialModels } from "./models.ts";
|
import { parseOfficialModels } from "./models.ts";
|
||||||
|
import { buildResponsesBody } from "cursor-byok:protocol/openai-responses";
|
||||||
import { codexProvider, isQuotaError } from "./provider.ts";
|
import { codexProvider, isQuotaError } from "./provider.ts";
|
||||||
import {
|
import {
|
||||||
accountIdentity,
|
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 () => {
|
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 token = jwt({ "https://api.openai.com/auth": { chatgpt_account_id: "acct-1" } });
|
||||||
const draft = await credentialDraft({
|
const draft = await credentialDraft({
|
||||||
|
|||||||
@@ -20,7 +20,6 @@ pub use descriptor::{
|
|||||||
};
|
};
|
||||||
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
pub use registry::{ImportResponse, OAuthBeginResponse, OAuthPollResponse, PluginRegistry};
|
||||||
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
pub use runtime::{PluginRuntime, PluginRuntimePhase, PluginRuntimeState, PluginRuntimeStatus};
|
||||||
pub(crate) use wire::llm_request as plugin_llm_request;
|
|
||||||
|
|
||||||
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
/// Windows 下阻止 Deno 子进程弹出控制台窗口(CREATE_NO_WINDOW)。
|
||||||
#[cfg(windows)]
|
#[cfg(windows)]
|
||||||
|
|||||||
@@ -20,8 +20,11 @@ use super::{
|
|||||||
worker::{PluginWorker, WorkerStreamItem},
|
worker::{PluginWorker, WorkerStreamItem},
|
||||||
};
|
};
|
||||||
use crate::{
|
use crate::{
|
||||||
model::ModelInvocation, provider::ModelEvent, provider::ProviderStream, store::Store, Error,
|
model::ModelInvocation,
|
||||||
Result,
|
provider::ProviderStream,
|
||||||
|
provider::{CallRecorder, ModelEvent},
|
||||||
|
store::Store,
|
||||||
|
Error, Result,
|
||||||
};
|
};
|
||||||
|
|
||||||
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
const OAUTH_SLOW_DOWN_STEP_MS: i64 = 5_000;
|
||||||
@@ -203,6 +206,7 @@ impl PluginRegistry {
|
|||||||
&self,
|
&self,
|
||||||
invocation: ModelInvocation,
|
invocation: ModelInvocation,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
|
recorder: CallRecorder,
|
||||||
) -> ProviderStream {
|
) -> ProviderStream {
|
||||||
let registry = self.clone();
|
let registry = self.clone();
|
||||||
Box::pin(try_stream! {
|
Box::pin(try_stream! {
|
||||||
@@ -232,7 +236,7 @@ impl PluginRegistry {
|
|||||||
"request": request,
|
"request": request,
|
||||||
});
|
});
|
||||||
let worker = registry.worker(&entry, &executable).await;
|
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() };
|
yield ModelEvent::Start { model_call_id: invocation.call_id.clone() };
|
||||||
while let Some(item) = items.recv().await {
|
while let Some(item) = items.recv().await {
|
||||||
match item {
|
match item {
|
||||||
|
|||||||
@@ -53,7 +53,17 @@ function replayItems(value: JsonValue): JsonValue[] {
|
|||||||
if (!Array.isArray(items)) {
|
if (!Array.isArray(items)) {
|
||||||
throw new Error("OpenAI Responses replay state is missing 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> {
|
export function buildResponsesBody(call: OpenAiResponsesCall): Record<string, JsonValue> {
|
||||||
|
|||||||
+269
-24
@@ -3,7 +3,10 @@ use std::{
|
|||||||
collections::{HashMap, HashSet},
|
collections::{HashMap, HashSet},
|
||||||
path::PathBuf,
|
path::PathBuf,
|
||||||
process::Stdio,
|
process::Stdio,
|
||||||
sync::Arc,
|
sync::{
|
||||||
|
atomic::{AtomicBool, Ordering},
|
||||||
|
Arc,
|
||||||
|
},
|
||||||
time::Duration,
|
time::Duration,
|
||||||
};
|
};
|
||||||
|
|
||||||
@@ -19,7 +22,7 @@ use super::{
|
|||||||
definition::{file_url, PluginDefinitionLoader},
|
definition::{file_url, PluginDefinitionLoader},
|
||||||
protocol::{HostMessage, WorkerMessage},
|
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 INVOCATION_TIMEOUT: Duration = Duration::from_secs(10 * 60);
|
||||||
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
const MAX_NETWORK_RESPONSE_BYTES: u64 = 16 * 1024 * 1024;
|
||||||
@@ -56,12 +59,29 @@ struct WorkerProcess {
|
|||||||
stdin: Arc<Mutex<ChildStdin>>,
|
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)]
|
#[derive(Clone)]
|
||||||
struct HostContext {
|
struct HostContext {
|
||||||
plugin_id: String,
|
plugin_id: String,
|
||||||
network_hosts: Arc<HashSet<String>>,
|
network_hosts: Arc<HashSet<String>>,
|
||||||
store: Store,
|
store: Store,
|
||||||
cancellations: Arc<Mutex<HashMap<String, CancellationToken>>>,
|
invocations: Arc<Mutex<HashMap<String, Arc<InvocationState>>>>,
|
||||||
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
streams: Arc<Mutex<HashMap<String, StreamLines>>>,
|
||||||
}
|
}
|
||||||
|
|
||||||
@@ -87,7 +107,7 @@ impl PluginWorker {
|
|||||||
.collect(),
|
.collect(),
|
||||||
),
|
),
|
||||||
store,
|
store,
|
||||||
cancellations: Arc::new(Mutex::new(HashMap::new())),
|
invocations: Arc::new(Mutex::new(HashMap::new())),
|
||||||
streams: Arc::new(Mutex::new(HashMap::new())),
|
streams: Arc::new(Mutex::new(HashMap::new())),
|
||||||
},
|
},
|
||||||
plugin_id,
|
plugin_id,
|
||||||
@@ -108,7 +128,9 @@ impl PluginWorker {
|
|||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
) -> Result<serde_json::Value> {
|
) -> 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 {
|
let result = tokio::time::timeout(INVOCATION_TIMEOUT, async {
|
||||||
while let Some(item) = items.recv().await {
|
while let Some(item) = items.recv().await {
|
||||||
if let WorkerStreamItem::Result(result) = item {
|
if let WorkerStreamItem::Result(result) = item {
|
||||||
@@ -137,15 +159,18 @@ impl PluginWorker {
|
|||||||
method: &str,
|
method: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
|
recorder: Option<CallRecorder>,
|
||||||
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
) -> Result<mpsc::UnboundedReceiver<WorkerStreamItem>> {
|
||||||
let id = uuid::Uuid::new_v4().to_string();
|
let id = uuid::Uuid::new_v4().to_string();
|
||||||
let request_cancellation = CancellationToken::new();
|
let request_cancellation = CancellationToken::new();
|
||||||
self.inner
|
self.inner.host.invocations.lock().await.insert(
|
||||||
.host
|
id.clone(),
|
||||||
.cancellations
|
Arc::new(InvocationState {
|
||||||
.lock()
|
cancellation: request_cancellation.clone(),
|
||||||
.await
|
recorder,
|
||||||
.insert(id.clone(), request_cancellation.clone());
|
recorder_claimed: AtomicBool::new(false),
|
||||||
|
}),
|
||||||
|
);
|
||||||
let (sender, receiver) = mpsc::unbounded_channel();
|
let (sender, receiver) = mpsc::unbounded_channel();
|
||||||
self.inner
|
self.inner
|
||||||
.pending
|
.pending
|
||||||
@@ -182,10 +207,10 @@ impl PluginWorker {
|
|||||||
}
|
}
|
||||||
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
let _ = sender.send(WorkerStreamItem::Result(Err(Error::Cancelled)));
|
||||||
inner.pending.lock().await.remove(&request_id);
|
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() => {
|
_ = 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) {
|
async fn cleanup(&self, id: &str) {
|
||||||
self.inner.pending.lock().await.remove(id);
|
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>>> {
|
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 {
|
impl HostContext {
|
||||||
async fn call(
|
async fn call(
|
||||||
&self,
|
&self,
|
||||||
@@ -421,7 +468,11 @@ impl HostContext {
|
|||||||
&self,
|
&self,
|
||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: &serde_json::Value,
|
params: &serde_json::Value,
|
||||||
) -> Result<(reqwest::RequestBuilder, CancellationToken)> {
|
) -> Result<(
|
||||||
|
reqwest::RequestBuilder,
|
||||||
|
CancellationToken,
|
||||||
|
Option<CallRecorder>,
|
||||||
|
)> {
|
||||||
let raw_url = required_string(params, "url")?;
|
let raw_url = required_string(params, "url")?;
|
||||||
let url = url::Url::parse(raw_url)
|
let url = url::Url::parse(raw_url)
|
||||||
.map_err(|error| Error::Config(format!("invalid plugin network URL: {error}")))?;
|
.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) {
|
if let Some(body) = params.get("body").and_then(serde_json::Value::as_str) {
|
||||||
request = request.body(body.to_owned());
|
request = request.body(body.to_owned());
|
||||||
}
|
}
|
||||||
let cancellation = self
|
let invocation = self.invocations.lock().await.get(request_id).cloned();
|
||||||
.cancellations
|
let cancellation = invocation
|
||||||
.lock()
|
.as_ref()
|
||||||
.await
|
.map(|state| state.cancellation.clone())
|
||||||
.get(request_id)
|
|
||||||
.cloned()
|
|
||||||
.unwrap_or_default();
|
.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(
|
async fn fetch(
|
||||||
@@ -478,13 +532,16 @@ impl HostContext {
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
) -> Result<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 request = request.timeout(Duration::from_secs(60));
|
||||||
let response = tokio::select! {
|
let response = tokio::select! {
|
||||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||||
response = request.send() => response?,
|
response = request.send() => response?,
|
||||||
};
|
};
|
||||||
let status = response.status().as_u16();
|
let status = response.status().as_u16();
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_headers(status).await?;
|
||||||
|
}
|
||||||
if response
|
if response
|
||||||
.content_length()
|
.content_length()
|
||||||
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
.is_some_and(|size| size > MAX_NETWORK_RESPONSE_BYTES)
|
||||||
@@ -503,6 +560,9 @@ impl HostContext {
|
|||||||
"plugin network response is larger than allowed".into(),
|
"plugin network response is larger than allowed".into(),
|
||||||
));
|
));
|
||||||
}
|
}
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_chunk(&body).await?;
|
||||||
|
}
|
||||||
Ok(
|
Ok(
|
||||||
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
serde_json::json!({ "status": status, "headers": headers, "body": String::from_utf8_lossy(&body) }),
|
||||||
)
|
)
|
||||||
@@ -514,12 +574,15 @@ impl HostContext {
|
|||||||
request_id: &str,
|
request_id: &str,
|
||||||
params: serde_json::Value,
|
params: serde_json::Value,
|
||||||
) -> Result<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! {
|
let response = tokio::select! {
|
||||||
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
_ = cancellation.cancelled() => return Err(Error::Cancelled),
|
||||||
response = request.send() => response?,
|
response = request.send() => response?,
|
||||||
};
|
};
|
||||||
let status = response.status().as_u16();
|
let status = response.status().as_u16();
|
||||||
|
if let Some(recorder) = &recorder {
|
||||||
|
recorder.response_headers(status).await?;
|
||||||
|
}
|
||||||
let headers = header_map(&response);
|
let headers = header_map(&response);
|
||||||
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
let (sender, receiver) = mpsc::channel::<Result<String>>(256);
|
||||||
tokio::spawn(async move {
|
tokio::spawn(async move {
|
||||||
@@ -552,6 +615,12 @@ impl HostContext {
|
|||||||
.await;
|
.await;
|
||||||
return;
|
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);
|
buffered.extend_from_slice(&chunk);
|
||||||
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
while let Some(position) = buffered.iter().position(|byte| *byte == b'\n') {
|
||||||
let mut line = buffered.drain(..=position).collect::<Vec<u8>>();
|
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)
|
.and_then(serde_json::Value::as_str)
|
||||||
.ok_or_else(|| Error::Protocol(format!("plugin host call requires string '{key}'")))
|
.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(|| {
|
.ok_or_else(|| {
|
||||||
Error::Protocol("OpenAI Responses replay state is missing items".into())
|
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);
|
push_responses_text(&mut input, &message.role, text);
|
||||||
for call in calls {
|
for call in calls {
|
||||||
@@ -437,6 +442,23 @@ fn responses_input(messages: &[ProjectedMessage]) -> Result<Vec<Value>> {
|
|||||||
Ok(input)
|
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<()> {
|
fn push_responses_parts(input: &mut Vec<Value>, role: &Role, parts: &[ContentPart]) -> Result<()> {
|
||||||
let text_type = if *role == Role::Assistant {
|
let text_type = if *role == Role::Assistant {
|
||||||
"output_text"
|
"output_text"
|
||||||
@@ -517,3 +539,47 @@ fn responses_usage(value: &Value) -> Usage {
|
|||||||
.and_then(Value::as_u64),
|
.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 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 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();
|
let guard = recorder.cancel_on_drop();
|
||||||
recorder.request(serde_json::json!({}), &crate::plugin::plugin_llm_request(&invocation)?).await?;
|
|
||||||
let mut routed = invocation.clone();
|
let mut routed = invocation.clone();
|
||||||
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
routed.request.model.display_name = Some(plan.model.display_name.clone());
|
||||||
if let Some(tokens) = plan.model.max_output_tokens {
|
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 {
|
let provider: Arc<dyn Provider> = Arc::new(NormalizedProvider::new(Arc::new(PluginModelProvider {
|
||||||
registry: plugins.clone(),
|
registry: plugins.clone(),
|
||||||
|
recorder: recorder.clone(),
|
||||||
})));
|
})));
|
||||||
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
(recorder, guard, provider.stream(routed, cancellation.clone()))
|
||||||
} else {
|
} else {
|
||||||
@@ -227,6 +227,7 @@ async fn finish_stream(recorder: &CallRecorder, cancellation: &CancellationToken
|
|||||||
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
/// 插件模型的 Provider 实现;对路由与规范化层完全等同于内置 Provider。
|
||||||
struct PluginModelProvider {
|
struct PluginModelProvider {
|
||||||
registry: PluginRegistry,
|
registry: PluginRegistry,
|
||||||
|
recorder: CallRecorder,
|
||||||
}
|
}
|
||||||
|
|
||||||
impl Provider for PluginModelProvider {
|
impl Provider for PluginModelProvider {
|
||||||
@@ -235,7 +236,8 @@ impl Provider for PluginModelProvider {
|
|||||||
invocation: ModelInvocation,
|
invocation: ModelInvocation,
|
||||||
cancellation: CancellationToken,
|
cancellation: CancellationToken,
|
||||||
) -> ProviderStream {
|
) -> ProviderStream {
|
||||||
self.registry.stream_model(invocation, cancellation)
|
self.registry
|
||||||
|
.stream_model(invocation, cancellation, self.recorder.clone())
|
||||||
}
|
}
|
||||||
}
|
}
|
||||||
|
|
||||||
|
|||||||
Reference in New Issue
Block a user