feat: NormalizedProvider

This commit is contained in:
leokun
2026-08-25 20:50:57 +08:00
parent 5f87357681
commit 081e1f50e2
8 changed files with 169 additions and 14 deletions
File diff suppressed because one or more lines are too long
+106 -1
View File
@@ -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
View File
@@ -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;
+28
View File
@@ -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)
}
}
+5 -4
View File
@@ -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)))
} }
+2 -1
View File
@@ -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,
+4 -1
View File
@@ -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);
+22 -6
View File
@@ -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