mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 12:13:05 +08:00
feat: NormalizedProvider
This commit is contained in:
@@ -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);
|
||||
|
||||
Reference in New Issue
Block a user