mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
fix: harden provider, tool, and desktop behavior
This commit is contained in:
@@ -1,3 +1,5 @@
|
||||
LOCAL_TAURI_SIGNING_KEY := $(CURDIR)/.tauri/cursor-byok.local.key
|
||||
|
||||
.PHONY: check dev-web dev-server dev-desktop build-web build-server build-desktop build-docker
|
||||
|
||||
check:
|
||||
@@ -21,8 +23,13 @@ build-web:
|
||||
build-server:
|
||||
cargo build --release --package cursor-server --bin cursor-server
|
||||
|
||||
build-desktop:
|
||||
npm --prefix apps/desktop run tauri:build
|
||||
$(LOCAL_TAURI_SIGNING_KEY):
|
||||
@install -d -m 700 "$(dir $@)"
|
||||
@apps/desktop/node_modules/.bin/tauri signer generate --ci --write-keys "$@" >/dev/null
|
||||
@chmod 600 "$@" "$@.pub"
|
||||
|
||||
build-desktop: $(LOCAL_TAURI_SIGNING_KEY)
|
||||
TAURI_SIGNING_PRIVATE_KEY="$(LOCAL_TAURI_SIGNING_KEY)" TAURI_SIGNING_PRIVATE_KEY_PASSWORD="" npm --prefix apps/desktop run tauri:build
|
||||
|
||||
build-docker:
|
||||
docker build --tag cursor-byok:local .
|
||||
|
||||
@@ -174,6 +174,7 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
|
||||
const shouldResetScrollRef = useRef(false)
|
||||
const scrollApiRef = useRef<ScrollAreaApi | null>(null)
|
||||
const scrollStateRef = useRef<ScrollAreaState | null>(null)
|
||||
const contentElementRef = useRef<HTMLDivElement | null>(null)
|
||||
const spacerRef = useRef<HTMLDivElement | null>(null)
|
||||
const [contentInsets, setContentInsets] = useState<ContentInsets>({
|
||||
top: 0,
|
||||
@@ -192,7 +193,8 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
|
||||
})
|
||||
const [, forceUpdate] = useState(0)
|
||||
|
||||
const setContentRef = useCallback((node: HTMLDivElement | null) => {
|
||||
const readContentInsets = useCallback(() => {
|
||||
const node = contentElementRef.current
|
||||
const styles = node ? getComputedStyle(node) : null
|
||||
const nextInsets = {
|
||||
top: styles ? Number.parseFloat(styles.paddingTop) || 0 : 0,
|
||||
@@ -205,6 +207,26 @@ export function VirtualList<TItem>(props: VirtualListProps<TItem>) {
|
||||
)
|
||||
}, [])
|
||||
|
||||
const setContentRef = useCallback((node: HTMLDivElement | null) => {
|
||||
contentElementRef.current = node
|
||||
readContentInsets()
|
||||
}, [readContentInsets])
|
||||
|
||||
useLayoutEffect(() => {
|
||||
const node = contentElementRef.current
|
||||
if (!node) return
|
||||
|
||||
readContentInsets()
|
||||
const resizeObserver = new ResizeObserver(readContentInsets)
|
||||
resizeObserver.observe(node)
|
||||
const frame = requestAnimationFrame(readContentInsets)
|
||||
|
||||
return () => {
|
||||
cancelAnimationFrame(frame)
|
||||
resizeObserver.disconnect()
|
||||
}
|
||||
}, [readContentInsets])
|
||||
|
||||
const contentInsetTop = contentInsets.top
|
||||
|
||||
if (!scrollStateRef.current) {
|
||||
|
||||
@@ -7,7 +7,7 @@ use crate::{
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
use super::{mcp_state, ReadImage, ToolCompletion};
|
||||
use super::{gate, mcp_state, ReadImage, ToolCompletion};
|
||||
use crate::cursor::tools::{
|
||||
edit,
|
||||
runtime::{ExecStage, PendingExec},
|
||||
@@ -18,6 +18,15 @@ pub(crate) fn from_exec(
|
||||
wire_result: &pb::exec_client_message::Message,
|
||||
) -> Result<ToolCompletion> {
|
||||
use pb::{exec_client_message::Message, tool_call::Tool};
|
||||
let mut gated_shell = matches!(
|
||||
wire_result,
|
||||
Message::ShellResult(_) | Message::MiniSweAgentBashResult(_)
|
||||
)
|
||||
.then(|| wire_result.clone());
|
||||
if let Some(message) = gated_shell.as_mut() {
|
||||
gate::exec_message(message);
|
||||
}
|
||||
let wire_result = gated_shell.as_ref().unwrap_or(wire_result);
|
||||
if let Message::McpStateExecResult(result) = wire_result {
|
||||
return mcp_state::complete(pending, result);
|
||||
}
|
||||
|
||||
@@ -0,0 +1,169 @@
|
||||
use crate::cursor::proto::agent::v1 as pb;
|
||||
|
||||
const KIB: usize = 1024;
|
||||
const SHELL_STREAM_LIMIT: usize = 16 * KIB;
|
||||
const SHELL_CONTENT_LIMIT: usize = 32 * KIB;
|
||||
|
||||
pub(super) fn model_content(tool: &pb::tool_call::Tool, content: &mut String) {
|
||||
if matches!(tool, pb::tool_call::Tool::ShellToolCall(_)) {
|
||||
*content = truncate_edges("Shell", content, SHELL_CONTENT_LIMIT);
|
||||
}
|
||||
}
|
||||
|
||||
pub(super) fn exec_message(message: &mut pb::exec_client_message::Message) {
|
||||
use pb::exec_client_message::Message;
|
||||
match message {
|
||||
Message::ShellResult(result) | Message::MiniSweAgentBashResult(result) => {
|
||||
gate_shell_result(result)
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn gate_shell_result(result: &mut pb::ShellResult) {
|
||||
use pb::shell_result::Result;
|
||||
match result.result.as_mut() {
|
||||
Some(Result::Success(success)) => {
|
||||
success.stdout = truncate_edges("Shell stdout", &success.stdout, SHELL_STREAM_LIMIT);
|
||||
success.stderr = truncate_edges("Shell stderr", &success.stderr, SHELL_STREAM_LIMIT);
|
||||
if let Some(interleaved) = success.interleaved_output.as_mut() {
|
||||
*interleaved =
|
||||
truncate_edges("Shell interleaved output", interleaved, SHELL_CONTENT_LIMIT);
|
||||
}
|
||||
}
|
||||
Some(Result::Failure(failure)) => {
|
||||
failure.stdout = truncate_edges("Shell stdout", &failure.stdout, SHELL_STREAM_LIMIT);
|
||||
failure.stderr = truncate_edges("Shell stderr", &failure.stderr, SHELL_STREAM_LIMIT);
|
||||
if let Some(interleaved) = failure.interleaved_output.as_mut() {
|
||||
*interleaved =
|
||||
truncate_edges("Shell interleaved output", interleaved, SHELL_CONTENT_LIMIT);
|
||||
}
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
|
||||
fn truncate_edges(tool_name: &str, content: &str, limit: usize) -> String {
|
||||
if content.len() <= limit {
|
||||
return content.to_string();
|
||||
}
|
||||
let original = content.len();
|
||||
let mut shown = limit;
|
||||
loop {
|
||||
let notice = format!(
|
||||
"\n\n[truncated: {tool_name} result exceeded {limit} bytes; omitted middle; showing {shown} of {original} bytes]\n\n"
|
||||
);
|
||||
let available = limit.saturating_sub(notice.len());
|
||||
let head = utf8_prefix(content, available / 2);
|
||||
let tail = utf8_suffix(content, available.saturating_sub(head.len()));
|
||||
let next_shown = head.len().saturating_add(tail.len());
|
||||
if next_shown == shown {
|
||||
return format!("{head}{notice}{tail}");
|
||||
}
|
||||
shown = next_shown;
|
||||
}
|
||||
}
|
||||
|
||||
fn utf8_prefix(value: &str, limit: usize) -> &str {
|
||||
let mut end = limit.min(value.len());
|
||||
while end > 0 && !value.is_char_boundary(end) {
|
||||
end -= 1;
|
||||
}
|
||||
&value[..end]
|
||||
}
|
||||
|
||||
fn utf8_suffix(value: &str, limit: usize) -> &str {
|
||||
let mut start = value.len().saturating_sub(limit);
|
||||
while start < value.len() && !value.is_char_boundary(start) {
|
||||
start += 1;
|
||||
}
|
||||
&value[start..]
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn shell_tool() -> pb::tool_call::Tool {
|
||||
pb::tool_call::Tool::ShellToolCall(pb::ShellToolCall::default())
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shell_output_keeps_both_ends_within_its_budget() {
|
||||
let mut content = format!("HEAD{}TAIL", " ".repeat(1024 * KIB));
|
||||
|
||||
model_content(&shell_tool(), &mut content);
|
||||
|
||||
assert!(content.len() <= SHELL_CONTENT_LIMIT);
|
||||
assert!(content.starts_with("HEAD"));
|
||||
assert!(content.ends_with("TAIL"));
|
||||
assert!(content.contains("omitted middle"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn non_shell_output_is_unchanged() {
|
||||
let mut content = "x".repeat(64 * KIB);
|
||||
let original = content.clone();
|
||||
|
||||
model_content(
|
||||
&pb::tool_call::Tool::ReadToolCall(pb::ReadToolCall::default()),
|
||||
&mut content,
|
||||
);
|
||||
|
||||
assert_eq!(content, original);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn shell_streams_are_limited_before_rendering() {
|
||||
let mut message = pb::exec_client_message::Message::ShellResult(pb::ShellResult {
|
||||
result: Some(pb::shell_result::Result::Success(pb::ShellSuccess {
|
||||
stdout: format!("HEAD{}TAIL", "x".repeat(64 * KIB)),
|
||||
stderr: format!("ERROR_HEAD{}ERROR_TAIL", "y".repeat(64 * KIB)),
|
||||
interleaved_output: Some(format!("START{}END", "z".repeat(64 * KIB))),
|
||||
..Default::default()
|
||||
})),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
exec_message(&mut message);
|
||||
|
||||
let pb::exec_client_message::Message::ShellResult(result) = message else {
|
||||
panic!("expected Shell result");
|
||||
};
|
||||
let Some(pb::shell_result::Result::Success(success)) = result.result else {
|
||||
panic!("expected Shell success");
|
||||
};
|
||||
assert!(success.stdout.len() <= SHELL_STREAM_LIMIT);
|
||||
assert!(success.stdout.starts_with("HEAD"));
|
||||
assert!(success.stdout.ends_with("TAIL"));
|
||||
assert!(success.stderr.len() <= SHELL_STREAM_LIMIT);
|
||||
assert!(success.stderr.starts_with("ERROR_HEAD"));
|
||||
assert!(success.stderr.ends_with("ERROR_TAIL"));
|
||||
assert!(success.interleaved_output.unwrap().len() <= SHELL_CONTENT_LIMIT);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn failed_shell_streams_are_limited() {
|
||||
let mut message = pb::exec_client_message::Message::ShellResult(pb::ShellResult {
|
||||
result: Some(pb::shell_result::Result::Failure(pb::ShellFailure {
|
||||
stdout: "x".repeat(64 * KIB),
|
||||
stderr: "y".repeat(64 * KIB),
|
||||
interleaved_output: Some("z".repeat(64 * KIB)),
|
||||
..Default::default()
|
||||
})),
|
||||
..Default::default()
|
||||
});
|
||||
|
||||
exec_message(&mut message);
|
||||
|
||||
let pb::exec_client_message::Message::ShellResult(result) = message else {
|
||||
panic!("expected Shell result");
|
||||
};
|
||||
let Some(pb::shell_result::Result::Failure(failure)) = result.result else {
|
||||
panic!("expected Shell failure");
|
||||
};
|
||||
assert!(failure.stdout.len() <= SHELL_STREAM_LIMIT);
|
||||
assert!(failure.stderr.len() <= SHELL_STREAM_LIMIT);
|
||||
assert!(failure.interleaved_output.unwrap().len() <= SHELL_CONTENT_LIMIT);
|
||||
}
|
||||
}
|
||||
@@ -1,5 +1,6 @@
|
||||
mod await_shell;
|
||||
mod exec;
|
||||
mod gate;
|
||||
mod interaction;
|
||||
mod local;
|
||||
mod mcp;
|
||||
@@ -87,9 +88,10 @@ impl ToolCompletion {
|
||||
pub(crate) fn new(
|
||||
call: &ToolCall,
|
||||
started_at_ms: u64,
|
||||
result: ToolResult,
|
||||
mut result: ToolResult,
|
||||
tool: pb::tool_call::Tool,
|
||||
) -> Self {
|
||||
gate::model_content(&tool, &mut result.content);
|
||||
Self {
|
||||
result,
|
||||
tool_call: pb::ToolCall {
|
||||
|
||||
@@ -138,7 +138,12 @@ pub fn normalize_base_url(value: &str) -> Result<String> {
|
||||
Ok(url.as_str().trim_end_matches('/').to_string())
|
||||
}
|
||||
|
||||
pub fn model_hash(base_url: &str, provider_type: ProviderType, model_id: &str) -> Result<String> {
|
||||
pub fn model_hash(
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
provider_type: ProviderType,
|
||||
model_id: &str,
|
||||
) -> Result<String> {
|
||||
let base_url = normalize_base_url(base_url)?;
|
||||
let model_id = model_id.trim();
|
||||
if model_id.is_empty() {
|
||||
@@ -147,6 +152,8 @@ pub fn model_hash(base_url: &str, provider_type: ProviderType, model_id: &str) -
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(base_url.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(api_key.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(provider_type.as_str().as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(model_id.as_bytes());
|
||||
@@ -212,24 +219,41 @@ mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn hash_uses_normalized_url_type_and_model_only() {
|
||||
fn hash_uses_normalized_url_key_type_and_model() {
|
||||
let first = model_hash(
|
||||
"HTTPS://Example.COM/v1/",
|
||||
"secret",
|
||||
ProviderType::OpenAiChat,
|
||||
"model-a",
|
||||
)
|
||||
.unwrap();
|
||||
let second = model_hash(
|
||||
"https://example.com/v1",
|
||||
"secret",
|
||||
ProviderType::OpenAiChat,
|
||||
"model-a",
|
||||
)
|
||||
.unwrap();
|
||||
assert_eq!(first, second);
|
||||
assert_eq!(first, "f246010a");
|
||||
assert_ne!(
|
||||
first,
|
||||
model_hash("https://example.com/v1", ProviderType::Anthropic, "model-a").unwrap()
|
||||
model_hash(
|
||||
"https://example.com/v1",
|
||||
"different-secret",
|
||||
ProviderType::OpenAiChat,
|
||||
"model-a",
|
||||
)
|
||||
.unwrap()
|
||||
);
|
||||
assert_ne!(
|
||||
first,
|
||||
model_hash(
|
||||
"https://example.com/v1",
|
||||
"secret",
|
||||
ProviderType::Anthropic,
|
||||
"model-a",
|
||||
)
|
||||
.unwrap()
|
||||
);
|
||||
}
|
||||
|
||||
|
||||
@@ -204,32 +204,6 @@ impl Provider for OpenAiResponsesProvider {
|
||||
}
|
||||
"response.completed" => {
|
||||
if let Some(usage) = value.pointer("/response/usage") { yield ModelEvent::Usage(responses_usage(usage)); }
|
||||
if let Some(output) = value.pointer("/response/output").and_then(Value::as_array) {
|
||||
for (index, item) in output.iter().enumerate() {
|
||||
match item.get("type").and_then(Value::as_str) {
|
||||
Some("reasoning") => {
|
||||
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
|
||||
if !reasoning_items.iter().any(|existing| existing.get("id") == item.get("id")) {
|
||||
reasoning_items.push(item.clone());
|
||||
}
|
||||
}
|
||||
Some("message") => {
|
||||
if let Some(final_text) = response_item_text(item) {
|
||||
for event in reconcile_response_text(&mut text_open, &mut text, &final_text) { yield event; }
|
||||
}
|
||||
}
|
||||
Some("function_call") => {
|
||||
saw_tool = true;
|
||||
let arguments = item
|
||||
.get("arguments")
|
||||
.and_then(Value::as_str)
|
||||
.map_or(ResponseToolArguments::None, ResponseToolArguments::Snapshot);
|
||||
for event in update_response_tool(index, item, arguments, true, &mut tools)? { yield event; }
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
}
|
||||
if thinking_open { thinking_open = false; yield ModelEvent::ThinkingEnd; }
|
||||
if text_open { text_open = false; yield ModelEvent::TextEnd; }
|
||||
for (index, tool) in tools.iter_mut().filter(|(_, tool)| tool.started && !tool.ended) {
|
||||
|
||||
@@ -36,7 +36,12 @@ impl Store {
|
||||
let mut hashes = Vec::with_capacity(models.len());
|
||||
let mut unique_hashes = HashSet::with_capacity(models.len());
|
||||
for model in models {
|
||||
let hash = model_hash(&base_url, model.endpoint_type, &model.model_id)?;
|
||||
let hash = model_hash(
|
||||
&base_url,
|
||||
provider.api_key.as_deref().unwrap_or_default(),
|
||||
model.endpoint_type,
|
||||
&model.model_id,
|
||||
)?;
|
||||
if !unique_hashes.insert(hash.clone()) {
|
||||
return Err(Error::Config(format!(
|
||||
"8-character model hash collision: {hash}"
|
||||
@@ -125,8 +130,8 @@ impl Store {
|
||||
let api_key = input.api_key.as_deref().unwrap_or(¤t.api_key);
|
||||
let custom_headers = merge_custom_headers(¤t.custom_headers, &input.custom_headers)?;
|
||||
let base_url = normalize_base_url(&input.base_url)?;
|
||||
let base_url_changed = base_url != current.endpoint.base_url;
|
||||
let models = if base_url_changed {
|
||||
let identity_changed = base_url != current.endpoint.base_url || api_key != current.api_key;
|
||||
let models = if identity_changed {
|
||||
sqlx::query("SELECT * FROM provider_models WHERE provider_id = ?")
|
||||
.bind(provider_id)
|
||||
.fetch_all(&self.pool)
|
||||
@@ -140,7 +145,7 @@ impl Store {
|
||||
let mut next_hashes = Vec::with_capacity(models.len());
|
||||
let mut unique_hashes = HashSet::with_capacity(models.len());
|
||||
for model in &models {
|
||||
let hash = model_hash(&base_url, model.endpoint_type, &model.model_id)?;
|
||||
let hash = model_hash(&base_url, api_key, model.endpoint_type, &model.model_id)?;
|
||||
if !unique_hashes.insert(hash.clone()) {
|
||||
return Err(Error::Config(format!(
|
||||
"8-character model hash collision: {hash}"
|
||||
@@ -271,6 +276,7 @@ impl Store {
|
||||
for input in inputs {
|
||||
let hash = model_hash(
|
||||
&provider.endpoint.base_url,
|
||||
&provider.api_key,
|
||||
input.endpoint_type,
|
||||
&input.model_id,
|
||||
)?;
|
||||
@@ -327,6 +333,7 @@ impl Store {
|
||||
.expect("model provider must exist");
|
||||
let next_hash = model_hash(
|
||||
&provider.endpoint.base_url,
|
||||
&provider.api_key,
|
||||
input.endpoint_type,
|
||||
&input.model_id,
|
||||
)?;
|
||||
@@ -680,6 +687,33 @@ mod tests {
|
||||
assert_eq!(store.provider_models(false).await.unwrap().len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn allows_same_endpoint_and_model_with_different_api_keys() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("credential-models.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let first_provider = provider();
|
||||
let mut second_provider = provider();
|
||||
second_provider.name = "Second".into();
|
||||
second_provider.api_key = Some("different-secret".into());
|
||||
|
||||
let (_, first_model) = store
|
||||
.create_provider_with_model(&first_provider, &model("model-a"))
|
||||
.await
|
||||
.unwrap();
|
||||
let (_, second_model) = store
|
||||
.create_provider_with_model(&second_provider, &model("model-a"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_ne!(first_model.model_hash, second_model.model_hash);
|
||||
assert_eq!(store.provider_models(false).await.unwrap().len(), 2);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn adds_multiple_models_to_existing_provider_atomically() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
@@ -750,6 +784,55 @@ mod tests {
|
||||
models[0].model_hash,
|
||||
model_hash(
|
||||
&updated_provider.base_url,
|
||||
input.api_key.as_deref().unwrap(),
|
||||
models[0].endpoint_type,
|
||||
&models[0].model_id,
|
||||
)
|
||||
.unwrap()
|
||||
);
|
||||
let detached: Option<String> =
|
||||
sqlx::query_scalar("SELECT model_hash FROM llm_calls WHERE call_id = ?")
|
||||
.bind("call-1")
|
||||
.fetch_one(store.pool())
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(detached, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn updating_provider_api_key_rehashes_its_models() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("provider-key-update.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let (created_provider, original) = store
|
||||
.create_provider_with_model(&provider(), &model("model-a"))
|
||||
.await
|
||||
.unwrap();
|
||||
insert_call(&store, &created_provider, &original).await;
|
||||
|
||||
let mut input = provider();
|
||||
input.api_key = Some("different-secret".into());
|
||||
store
|
||||
.update_provider(created_provider.provider_id, &input)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert!(store
|
||||
.provider_model(&original.model_hash)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_none());
|
||||
let models = store.provider_models(false).await.unwrap();
|
||||
assert_eq!(models.len(), 1);
|
||||
assert_eq!(
|
||||
models[0].model_hash,
|
||||
model_hash(
|
||||
&created_provider.base_url,
|
||||
"different-secret",
|
||||
models[0].endpoint_type,
|
||||
&models[0].model_id,
|
||||
)
|
||||
|
||||
@@ -77,7 +77,7 @@ async fn provider_secret_is_write_only_and_model_hash_is_stable() {
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert_eq!(model.model_hash, "f246010a");
|
||||
assert_eq!(model.model_hash, "bab5019a");
|
||||
assert!(model.supports_image_generation);
|
||||
}
|
||||
|
||||
|
||||
@@ -148,6 +148,35 @@ async fn duplicate_usage_is_rejected_instead_of_guessing_which_total_is_final()
|
||||
assert!(matches!(failure.failure, RunFailure::Protocol(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn duplicate_tool_call_ids_are_rejected_across_distinct_indexes() {
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(8);
|
||||
let failure = consume_model_cycle(
|
||||
provider_stream(vec![
|
||||
ModelEvent::Start {
|
||||
model_call_id: "model-call".into(),
|
||||
},
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 0,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
ModelEvent::ToolCallEnd { index: 0 },
|
||||
ModelEvent::ToolCallStart {
|
||||
index: 1,
|
||||
call_id: "call-1".into(),
|
||||
name: "Read".into(),
|
||||
},
|
||||
]),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap_err();
|
||||
|
||||
assert!(matches!(failure.failure, RunFailure::Protocol(_)));
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_chat_raw_stream_and_request_projection_match_the_endpoint() {
|
||||
let (base_url, mut requests, server) = fixture_server(
|
||||
@@ -457,12 +486,15 @@ async fn openai_responses_preserves_delta_that_repeats_the_streamed_suffix() {
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn openai_responses_completed_object_recovers_missing_item_events() {
|
||||
async fn openai_responses_completed_snapshot_does_not_reindex_streamed_tool() {
|
||||
let (base_url, _requests, server) = fixture_server(
|
||||
"/v1/responses",
|
||||
concat!(
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":0,\"item\":{\"type\":\"reasoning\",\"id\":\"reasoning-1\",\"encrypted_content\":\"opaque\"}}\n\n",
|
||||
"data: {\"type\":\"response.output_item.added\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\"}}\n\n",
|
||||
"data: {\"type\":\"response.function_call_arguments.done\",\"output_index\":1,\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}\n\n",
|
||||
"data: {\"type\":\"response.output_item.done\",\"output_index\":1,\"item\":{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}}\n\n",
|
||||
"data: {\"type\":\"response.completed\",\"response\":{\"output\":[",
|
||||
"{\"type\":\"message\",\"content\":[{\"type\":\"output_text\",\"text\":\"ok\"}]},",
|
||||
"{\"type\":\"function_call\",\"call_id\":\"call-1\",\"name\":\"Read\",\"arguments\":\"{\\\"path\\\":\\\"a\\\"}\"}",
|
||||
"]}}\n\n",
|
||||
),
|
||||
@@ -472,18 +504,21 @@ async fn openai_responses_completed_object_recovers_missing_item_events() {
|
||||
reqwest::Client::new(),
|
||||
config(ProviderKind::OpenAiResponses, base_url, None),
|
||||
);
|
||||
let (sender, _receiver) = tokio::sync::mpsc::channel(32);
|
||||
|
||||
let events = collect(provider.stream(invocation(), CancellationToken::new())).await;
|
||||
let cycle = consume_model_cycle(
|
||||
provider.stream(invocation(), CancellationToken::new()),
|
||||
&sender,
|
||||
&CancellationToken::new(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
server.abort();
|
||||
|
||||
assert!(events
|
||||
.iter()
|
||||
.any(|event| matches!(event, ModelEvent::TextDelta(text) if text == "ok")));
|
||||
assert!(events.iter().any(|event| matches!(event, ModelEvent::ToolCallStart { call_id, name, .. } if call_id == "call-1" && name == "Read")));
|
||||
assert_eq!(
|
||||
events.last(),
|
||||
Some(&ModelEvent::Done(FinishReason::ToolUse))
|
||||
);
|
||||
assert_eq!(cycle.calls.len(), 1);
|
||||
assert_eq!(cycle.calls[0].index, 1);
|
||||
assert_eq!(cycle.calls[0].call_id, "call-1");
|
||||
assert_eq!(cycle.calls[0].arguments["path"], "a");
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
|
||||
Reference in New Issue
Block a user