mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +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 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)]
|
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||||
pub struct PromptSpec {
|
pub struct PromptSpec {
|
||||||
@@ -23,3 +25,106 @@ pub struct ModelInvocation {
|
|||||||
pub provider_call_index: u64,
|
pub provider_call_index: u64,
|
||||||
pub request: ModelRequest,
|
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 anthropic;
|
||||||
mod event;
|
mod event;
|
||||||
|
mod normalize;
|
||||||
mod openai_chat;
|
mod openai_chat;
|
||||||
mod openai_responses;
|
mod openai_responses;
|
||||||
mod recorder;
|
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::{
|
use super::{
|
||||||
AnthropicProvider, CallRecorder, OpenAiChatProvider, OpenAiResponsesProvider, Provider,
|
normalize::NormalizedProvider, AnthropicProvider, CallRecorder, OpenAiChatProvider,
|
||||||
ProviderStream,
|
OpenAiResponsesProvider, Provider, ProviderStream,
|
||||||
};
|
};
|
||||||
|
|
||||||
pub struct ProviderRouter {
|
pub struct ProviderRouter {
|
||||||
@@ -160,7 +160,7 @@ fn build_inner(
|
|||||||
.timeout(config.request_timeout)
|
.timeout(config.request_timeout)
|
||||||
.build()?,
|
.build()?,
|
||||||
};
|
};
|
||||||
Ok(match config.kind {
|
let provider: Arc<dyn Provider> = match config.kind {
|
||||||
ProviderKind::OpenAiChat => {
|
ProviderKind::OpenAiChat => {
|
||||||
Arc::new(OpenAiChatProvider::new(client, config.clone()).with_recorder(recorder))
|
Arc::new(OpenAiChatProvider::new(client, config.clone()).with_recorder(recorder))
|
||||||
}
|
}
|
||||||
@@ -170,5 +170,6 @@ fn build_inner(
|
|||||||
ProviderKind::Anthropic => {
|
ProviderKind::Anthropic => {
|
||||||
Arc::new(AnthropicProvider::new(client, config.clone()).with_recorder(recorder))
|
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,
|
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(
|
return Err(failure(
|
||||||
RunFailure::Protocol("finish reason and tool calls disagree".into()),
|
RunFailure::Protocol("finish reason and tool calls disagree".into()),
|
||||||
text,
|
text,
|
||||||
|
|||||||
@@ -908,7 +908,10 @@ mod tests {
|
|||||||
assert_eq!(listed[0].api_key.as_deref(), Some("secret"));
|
assert_eq!(listed[0].api_key.as_deref(), Some("secret"));
|
||||||
assert!(listed[0].has_api_key);
|
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();
|
let empty = store.create_provider(&without_key).await.unwrap();
|
||||||
assert_eq!(empty.api_key, None);
|
assert_eq!(empty.api_key, None);
|
||||||
assert!(!empty.has_api_key);
|
assert!(!empty.has_api_key);
|
||||||
|
|||||||
@@ -172,7 +172,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
"SembleSearch",
|
"SembleSearch",
|
||||||
"SembleFindRelated",
|
"SembleFindRelated",
|
||||||
],
|
],
|
||||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
|
||||||
);
|
);
|
||||||
assert_mode(
|
assert_mode(
|
||||||
&assets,
|
&assets,
|
||||||
@@ -194,7 +194,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
"SembleSearch",
|
"SembleSearch",
|
||||||
"SembleFindRelated",
|
"SembleFindRelated",
|
||||||
],
|
],
|
||||||
"235a2a9a7785844eb5186f1c8f2294a36a04bbf103887d05ec386f8c7cc52abc",
|
"e2eb8a1ebd70d53b1b2eb6bedabdce62ff070a05a6168216013d0a1144ed8bb5",
|
||||||
);
|
);
|
||||||
assert_mode(
|
assert_mode(
|
||||||
&assets,
|
&assets,
|
||||||
@@ -218,7 +218,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
"SembleSearch",
|
"SembleSearch",
|
||||||
"SembleFindRelated",
|
"SembleFindRelated",
|
||||||
],
|
],
|
||||||
"cefa1800d7440611b6c3e922fa59f1fe262eca4cee52ea71043fe1f29e52c659",
|
"ec10becac85819cda321298762892852194c78601db66cc0b4ce74bc1213e29e",
|
||||||
);
|
);
|
||||||
assert_mode(
|
assert_mode(
|
||||||
&assets,
|
&assets,
|
||||||
@@ -244,7 +244,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
"SembleSearch",
|
"SembleSearch",
|
||||||
"SembleFindRelated",
|
"SembleFindRelated",
|
||||||
],
|
],
|
||||||
"04c5fb238eb3695936ceed610b481caf7507f934efb88cb7130a8756e12959e3",
|
"25f7b559941baabfc9b1046455b04ca812fc41a6878ad55a43d83f0bd18cd92f",
|
||||||
);
|
);
|
||||||
assert_mode(
|
assert_mode(
|
||||||
&assets,
|
&assets,
|
||||||
@@ -272,7 +272,7 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
"SembleSearch",
|
"SembleSearch",
|
||||||
"SembleFindRelated",
|
"SembleFindRelated",
|
||||||
],
|
],
|
||||||
"f88c55fdbb53be377e64cc6280ebb23c75e2d244b5c3463ed18e90752c3b7ff5",
|
"48c8e0fe825f9c2450307ca5e70cde7077c4282c135b2cd15338bd4bd0c43636",
|
||||||
);
|
);
|
||||||
assert_mode(
|
assert_mode(
|
||||||
&assets,
|
&assets,
|
||||||
@@ -282,8 +282,24 @@ fn every_prompt_mode_loads_the_captured_tool_set() {
|
|||||||
);
|
);
|
||||||
assert_eq!(
|
assert_eq!(
|
||||||
schema_digest(&assets.mode(Mode::Agent).tools),
|
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
|
let shell = assets
|
||||||
.mode(Mode::Agent)
|
.mode(Mode::Agent)
|
||||||
.tools
|
.tools
|
||||||
|
|||||||
Reference in New Issue
Block a user