mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-03 18:23:51 +08:00
feat: NormalizedProvider
This commit is contained in:
File diff suppressed because one or more lines are too long
@@ -1,6 +1,8 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{ModelSpec, ProjectedMessage, ToolDefinition};
|
||||
use super::{ModelSpec, ProjectedContent, ProjectedMessage, ToolDefinition};
|
||||
|
||||
const PROVIDER_TOOL_CALL_ID_MAX_CHARS: usize = 64;
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct PromptSpec {
|
||||
@@ -23,3 +25,106 @@ pub struct ModelInvocation {
|
||||
pub provider_call_index: u64,
|
||||
pub request: ModelRequest,
|
||||
}
|
||||
|
||||
pub(crate) fn normalize_provider_tool_call_ids(history: &mut [ProjectedMessage]) {
|
||||
for message in history {
|
||||
match &mut message.content {
|
||||
ProjectedContent::Assistant { calls, .. } => {
|
||||
for call in calls {
|
||||
truncate_tool_call_id(&mut call.call_id);
|
||||
}
|
||||
}
|
||||
ProjectedContent::ToolResult(result) => {
|
||||
truncate_tool_call_id(&mut result.call_id);
|
||||
}
|
||||
ProjectedContent::Parts(_) => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
fn truncate_tool_call_id(call_id: &mut String) {
|
||||
if let Some((end, _)) = call_id.char_indices().nth(PROVIDER_TOOL_CALL_ID_MAX_CHARS) {
|
||||
call_id.truncate(end);
|
||||
}
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::normalize_provider_tool_call_ids;
|
||||
use crate::model::{
|
||||
ProjectedContent, ProjectedMessage, Role, ToolCallContent, ToolResultContent,
|
||||
};
|
||||
|
||||
#[test]
|
||||
fn provider_tool_call_ids_are_truncated_once_for_every_provider() {
|
||||
let call_id = format!("cursor-tool-call:{}", "x".repeat(68));
|
||||
assert_eq!(call_id.len(), 85);
|
||||
let expected = call_id[..64].to_string();
|
||||
let mut history = vec![
|
||||
ProjectedMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: String::new(),
|
||||
replay_state: None,
|
||||
calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: call_id.clone(),
|
||||
name: "Shell".into(),
|
||||
arguments: serde_json::json!({}),
|
||||
}],
|
||||
},
|
||||
},
|
||||
ProjectedMessage {
|
||||
message_id: "result".into(),
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(ToolResultContent {
|
||||
call_id,
|
||||
name: "Shell".into(),
|
||||
content: "done".into(),
|
||||
is_error: false,
|
||||
image: None,
|
||||
provider_parts: Vec::new(),
|
||||
}),
|
||||
},
|
||||
];
|
||||
|
||||
normalize_provider_tool_call_ids(&mut history);
|
||||
|
||||
let ProjectedContent::Assistant { calls, .. } = &history[0].content else {
|
||||
panic!("expected assistant message");
|
||||
};
|
||||
let ProjectedContent::ToolResult(result) = &history[1].content else {
|
||||
panic!("expected tool result");
|
||||
};
|
||||
assert_eq!(calls[0].call_id, expected);
|
||||
assert_eq!(result.call_id, expected);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_tool_call_id_truncation_counts_unicode_characters() {
|
||||
let mut history = vec![ProjectedMessage {
|
||||
message_id: "assistant".into(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text: String::new(),
|
||||
thinking: String::new(),
|
||||
replay_state: None,
|
||||
calls: vec![ToolCallContent {
|
||||
index: 0,
|
||||
call_id: format!("{}界y", "x".repeat(63)),
|
||||
name: "Read".into(),
|
||||
arguments: serde_json::json!({}),
|
||||
}],
|
||||
},
|
||||
}];
|
||||
|
||||
normalize_provider_tool_call_ids(&mut history);
|
||||
|
||||
let ProjectedContent::Assistant { calls, .. } = &history[0].content else {
|
||||
panic!("expected assistant message");
|
||||
};
|
||||
assert_eq!(calls[0].call_id, format!("{}界", "x".repeat(63)));
|
||||
}
|
||||
}
|
||||
|
||||
@@ -1,5 +1,6 @@
|
||||
mod anthropic;
|
||||
mod event;
|
||||
mod normalize;
|
||||
mod openai_chat;
|
||||
mod openai_responses;
|
||||
mod recorder;
|
||||
|
||||
@@ -0,0 +1,28 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::model::{normalize_provider_tool_call_ids, ModelInvocation};
|
||||
|
||||
use super::{Provider, ProviderStream};
|
||||
|
||||
pub(super) struct NormalizedProvider {
|
||||
inner: Arc<dyn Provider>,
|
||||
}
|
||||
|
||||
impl NormalizedProvider {
|
||||
pub(super) fn new(inner: Arc<dyn Provider>) -> Self {
|
||||
Self { inner }
|
||||
}
|
||||
}
|
||||
|
||||
impl Provider for NormalizedProvider {
|
||||
fn stream(
|
||||
&self,
|
||||
mut invocation: ModelInvocation,
|
||||
cancellation: CancellationToken,
|
||||
) -> ProviderStream {
|
||||
normalize_provider_tool_call_ids(&mut invocation.request.history);
|
||||
self.inner.stream(invocation, cancellation)
|
||||
}
|
||||
}
|
||||
@@ -12,8 +12,8 @@ use crate::{
|
||||
};
|
||||
|
||||
use super::{
|
||||
AnthropicProvider, CallRecorder, OpenAiChatProvider, OpenAiResponsesProvider, Provider,
|
||||
ProviderStream,
|
||||
normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider,
|
||||
OpenAiResponsesProvider, Provider, ProviderStream,
|
||||
};
|
||||
|
||||
pub struct ProviderRouter {
|
||||
@@ -160,7 +160,7 @@ fn build_inner(
|
||||
.timeout(config.request_timeout)
|
||||
.build()?,
|
||||
};
|
||||
Ok(match config.kind {
|
||||
let provider: Arc<dyn Provider> = match config.kind {
|
||||
ProviderKind::OpenAiChat => {
|
||||
Arc::new(OpenAiChatProvider::new(client, config.clone()).with_recorder(recorder))
|
||||
}
|
||||
@@ -170,5 +170,6 @@ fn build_inner(
|
||||
ProviderKind::Anthropic => {
|
||||
Arc::new(AnthropicProvider::new(client, config.clone()).with_recorder(recorder))
|
||||
}
|
||||
})
|
||||
};
|
||||
Ok(Arc::new(NormalizedProvider::new(provider)))
|
||||
}
|
||||
|
||||
@@ -308,7 +308,8 @@ pub async fn consume_model_cycle(
|
||||
usage,
|
||||
));
|
||||
}
|
||||
if matches!(finish_reason, FinishReason::ToolUse) != !calls.is_empty() {
|
||||
let has_tool_calls = !calls.is_empty();
|
||||
if matches!(finish_reason, FinishReason::ToolUse) != has_tool_calls {
|
||||
return Err(failure(
|
||||
RunFailure::Protocol("finish reason and tool calls disagree".into()),
|
||||
text,
|
||||
|
||||
@@ -908,7 +908,10 @@ mod tests {
|
||||
assert_eq!(listed[0].api_key.as_deref(), Some("secret"));
|
||||
assert!(listed[0].has_api_key);
|
||||
|
||||
let without_key = ProviderEndpointInput { api_key: None, ..provider() };
|
||||
let without_key = ProviderEndpointInput {
|
||||
api_key: None,
|
||||
..provider()
|
||||
};
|
||||
let empty = store.create_provider(&without_key).await.unwrap();
|
||||
assert_eq!(empty.api_key, None);
|
||||
assert!(!empty.has_api_key);
|
||||
|
||||
@@ -172,7 +172,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
||||
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
@@ -194,7 +194,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"235a2a9a7785844eb5186f1c8f2294a36a04bbf103887d05ec386f8c7cc52abc",
|
||||
"e2eb8a1ebd70d53b1b2eb6bedabdce62ff070a05a6168216013d0a1144ed8bb5",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
@@ -218,7 +218,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
||||
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
@@ -244,7 +244,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"04c5fb238eb3695936ceed610b481caf7507f934efb88cb7130a8756e12959e3",
|
||||
"25f7b559941baabfc9b1046455b04ca812fc41a6878ad55a43d83f0bd18cd92f",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
@@ -272,7 +272,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
"SembleSearch",
|
||||
"SembleFindRelated",
|
||||
],
|
||||
"f88c55fdbb53be377e64cc6280ebb23c75e2d244b5c3463ed18e90752c3b7ff5",
|
||||
"48c8e0fe825f9c2450307ca5e70cde7077c4282c135b2cd15338bd4bd0c43636",
|
||||
);
|
||||
assert_mode(
|
||||
&assets,
|
||||
@@ -282,8 +282,24 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
||||
);
|
||||
assert_eq!(
|
||||
schema_digest(&assets.mode(Mode::Agent).tools),
|
||||
"4324c36fa047fbe4c93a5d5f0b736c559a942e097266c5bae057804289f8b359"
|
||||
"e53a72c1d131ff3f65c619799232440b064e90e99f5d3fcceb63e32598d3a0fc"
|
||||
);
|
||||
let task = assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
.iter()
|
||||
.find(|tool| tool.name == "Task")
|
||||
.unwrap();
|
||||
assert!(task.description.contains(
|
||||
"When the user does not specify a number, launch at most three subagents in a single response. If the user explicitly requests more, you may launch the requested number."
|
||||
));
|
||||
assert!(task.description.contains(
|
||||
"If the user explicitly requests parallel subagents, follow the number requested by the user."
|
||||
));
|
||||
assert!(!task
|
||||
.description
|
||||
.chars()
|
||||
.any(|character| ('\u{4e00}'..='\u{9fff}').contains(&character)));
|
||||
let shell = assets
|
||||
.mode(Mode::Agent)
|
||||
.tools
|
||||
|
||||
Reference in New Issue
Block a user