mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 04:07:36 +08:00
fix: adjust secondary button height and enhance MultiCombobox input behavior
This commit is contained in:
+154
-49
@@ -18,7 +18,7 @@ use crate::{
|
||||
ContentPart, CursorRunTraceArtifact, CursorRunTraceSummary, LlmCallRequest,
|
||||
LlmCallResponseChunk, LlmCallSummary, ModelInvocation, ModelRequest, ModelSpec, Overview,
|
||||
ProjectedContent, ProjectedMessage, PromptSpec, ProviderEndpoint, ProviderEndpointInput,
|
||||
ProviderEndpointSecret, ProviderModel, ProviderModelInput, ProviderType, Role,
|
||||
ProviderModel, ProviderModelInput, ProviderType, Role,
|
||||
},
|
||||
provider::{ModelEvent, Provider},
|
||||
store::{
|
||||
@@ -349,34 +349,15 @@ impl ControlService {
|
||||
|
||||
pub async fn discover_input(&self, input: &ProviderEndpointInput) -> Result<DiscoveredModels> {
|
||||
let client = crate::network::client(&self.store).await?;
|
||||
let endpoint = ProviderEndpoint {
|
||||
provider_id: 0,
|
||||
name: input.name.clone(),
|
||||
provider_type: input.provider_type,
|
||||
base_url: crate::model::normalize_base_url(&input.base_url)?,
|
||||
has_api_key: input
|
||||
.api_key
|
||||
.as_deref()
|
||||
.is_some_and(|value| !value.is_empty()),
|
||||
custom_headers: input.custom_headers.clone(),
|
||||
extra_params: input.extra_params.clone(),
|
||||
created_at_ms: 0,
|
||||
updated_at_ms: 0,
|
||||
};
|
||||
let secret = ProviderEndpointSecret {
|
||||
endpoint,
|
||||
api_key: input.api_key.clone().unwrap_or_default(),
|
||||
custom_headers: input.custom_headers.clone(),
|
||||
};
|
||||
let mut models = match input.provider_type {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
|
||||
openai_models(&client, &secret).await?
|
||||
}
|
||||
ProviderType::Anthropic => anthropic_models(&client, &secret).await?,
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
Ok(DiscoveredModels { models })
|
||||
let base_url = crate::model::normalize_base_url(&input.base_url)?;
|
||||
discover_provider_models(
|
||||
&client,
|
||||
input.provider_type,
|
||||
&base_url,
|
||||
input.api_key.as_deref().unwrap_or_default(),
|
||||
&input.custom_headers,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn discover_models(&self, provider_id: i64) -> Result<DiscoveredModels> {
|
||||
@@ -386,15 +367,14 @@ impl ControlService {
|
||||
.provider(provider_id)
|
||||
.await?
|
||||
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
|
||||
let mut models = match provider.endpoint.provider_type {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
|
||||
openai_models(&client, &provider).await?
|
||||
}
|
||||
ProviderType::Anthropic => anthropic_models(&client, &provider).await?,
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
Ok(DiscoveredModels { models })
|
||||
discover_provider_models(
|
||||
&client,
|
||||
provider.endpoint.provider_type,
|
||||
&provider.endpoint.base_url,
|
||||
provider.endpoint.api_key.as_deref().unwrap_or_default(),
|
||||
&provider.custom_headers,
|
||||
)
|
||||
.await
|
||||
}
|
||||
|
||||
pub async fn calls(&self, limit: i64) -> Result<Vec<CallSummary>> {
|
||||
@@ -606,15 +586,51 @@ fn readable_utf8(data: &[u8]) -> Option<&str> {
|
||||
.then_some(value)
|
||||
}
|
||||
|
||||
async fn discover_provider_models(
|
||||
client: &reqwest::Client,
|
||||
provider_type: ProviderType,
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<DiscoveredModels> {
|
||||
let mut models = match provider_type {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => {
|
||||
openai_models(client, base_url, api_key, custom_headers).await?
|
||||
}
|
||||
ProviderType::Anthropic => {
|
||||
anthropic_models(client, base_url, api_key, custom_headers).await?
|
||||
}
|
||||
};
|
||||
models.sort();
|
||||
models.dedup();
|
||||
Ok(DiscoveredModels { models })
|
||||
}
|
||||
|
||||
fn model_discovery_url(base_url: &str) -> Result<Url> {
|
||||
let mut url = Url::parse(base_url)
|
||||
.map_err(|error| Error::Config(format!("invalid provider base URL: {error}")))?;
|
||||
if url.host_str().is_none() {
|
||||
return Err(Error::Config(
|
||||
"provider base URL must contain a host".into(),
|
||||
));
|
||||
}
|
||||
url.set_path("/v1/models");
|
||||
url.set_query(None);
|
||||
url.set_fragment(None);
|
||||
Ok(url)
|
||||
}
|
||||
|
||||
async fn openai_models(
|
||||
client: &reqwest::Client,
|
||||
provider: &ProviderEndpointSecret,
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut request = client.get(format!("{}/models", provider.endpoint.base_url));
|
||||
if !provider.api_key.is_empty() {
|
||||
request = request.bearer_auth(&provider.api_key);
|
||||
let mut request = client.get(model_discovery_url(base_url)?);
|
||||
if !api_key.is_empty() {
|
||||
request = request.bearer_auth(api_key);
|
||||
}
|
||||
let response = apply_custom_headers(request, &provider.custom_headers)?
|
||||
let response = apply_discovery_headers(request, custom_headers)?
|
||||
.send()
|
||||
.await?;
|
||||
let status = response.status();
|
||||
@@ -629,22 +645,24 @@ async fn openai_models(
|
||||
|
||||
async fn anthropic_models(
|
||||
client: &reqwest::Client,
|
||||
provider: &ProviderEndpointSecret,
|
||||
base_url: &str,
|
||||
api_key: &str,
|
||||
custom_headers: &serde_json::Value,
|
||||
) -> Result<Vec<String>> {
|
||||
let mut after_id = None::<String>;
|
||||
let mut found = BTreeSet::new();
|
||||
loop {
|
||||
let mut request = client
|
||||
.get(format!("{}/models", provider.endpoint.base_url))
|
||||
.get(model_discovery_url(base_url)?)
|
||||
.query(&[("limit", "100")])
|
||||
.header("anthropic-version", "2023-06-01");
|
||||
if !provider.api_key.is_empty() {
|
||||
request = request.header("x-api-key", &provider.api_key);
|
||||
if !api_key.is_empty() {
|
||||
request = request.header("x-api-key", api_key);
|
||||
}
|
||||
if let Some(after_id) = &after_id {
|
||||
request = request.query(&[("after_id", after_id)]);
|
||||
}
|
||||
let response = apply_custom_headers(request, &provider.custom_headers)?
|
||||
let response = apply_discovery_headers(request, custom_headers)?
|
||||
.send()
|
||||
.await?;
|
||||
let status = response.status();
|
||||
@@ -699,7 +717,7 @@ fn estimate_output_tokens(output: &str) -> u64 {
|
||||
}
|
||||
}
|
||||
|
||||
fn apply_custom_headers(
|
||||
fn apply_discovery_headers(
|
||||
mut request: reqwest::RequestBuilder,
|
||||
headers: &serde_json::Value,
|
||||
) -> Result<reqwest::RequestBuilder> {
|
||||
@@ -707,6 +725,9 @@ fn apply_custom_headers(
|
||||
.as_object()
|
||||
.ok_or_else(|| Error::Config("custom headers must be an object".into()))?;
|
||||
for (name, value) in object {
|
||||
if name.eq_ignore_ascii_case("user-agent") {
|
||||
continue;
|
||||
}
|
||||
let value = value
|
||||
.as_str()
|
||||
.ok_or_else(|| Error::Config(format!("custom header {name} must be a string")))?;
|
||||
@@ -838,4 +859,88 @@ mod tests {
|
||||
assert_eq!(super::estimate_output_tokens("1 2 3"), 3);
|
||||
assert_eq!(super::estimate_output_tokens(""), 0);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn model_discovery_url_uses_only_the_provider_origin() {
|
||||
assert_eq!(
|
||||
super::model_discovery_url("https://example.com:8443/arbitrary/v1/chat/completions")
|
||||
.unwrap()
|
||||
.as_str(),
|
||||
"https://example.com:8443/v1/models"
|
||||
);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn model_discovery_does_not_inherit_user_agent_or_request_body_settings() {
|
||||
type CapturedRequest = (
|
||||
axum::http::Method,
|
||||
axum::http::Uri,
|
||||
axum::http::HeaderMap,
|
||||
bytes::Bytes,
|
||||
);
|
||||
|
||||
async fn models(
|
||||
axum::extract::State(sender): axum::extract::State<
|
||||
tokio::sync::mpsc::UnboundedSender<CapturedRequest>,
|
||||
>,
|
||||
request: axum::extract::Request,
|
||||
) -> axum::Json<serde_json::Value> {
|
||||
let (parts, body) = request.into_parts();
|
||||
let body = axum::body::to_bytes(body, usize::MAX).await.unwrap();
|
||||
sender
|
||||
.send((parts.method, parts.uri, parts.headers, body))
|
||||
.unwrap();
|
||||
axum::Json(serde_json::json!({ "data": [{ "id": "model-a" }] }))
|
||||
}
|
||||
|
||||
let (sender, mut requests) = tokio::sync::mpsc::unbounded_channel();
|
||||
let app = axum::Router::new()
|
||||
.route("/v1/models", axum::routing::get(models))
|
||||
.with_state(sender);
|
||||
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
||||
let address = listener.local_addr().unwrap();
|
||||
let server = tokio::spawn(async move { axum::serve(listener, app).await.unwrap() });
|
||||
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("discovery.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let service = ControlService::new(
|
||||
store,
|
||||
Arc::new(TestProvider {
|
||||
invocation: Arc::new(Mutex::new(None)),
|
||||
}),
|
||||
)
|
||||
.unwrap();
|
||||
let result = service
|
||||
.discover_input(&ProviderEndpointInput {
|
||||
name: "Test".into(),
|
||||
provider_type: ProviderType::OpenAiResponses,
|
||||
base_url: format!("http://{address}/custom/responses"),
|
||||
api_key: Some("secret".into()),
|
||||
custom_headers: serde_json::json!({
|
||||
"uSeR-aGeNt": "inherited-user-agent",
|
||||
"x-tenant": "tenant-a"
|
||||
}),
|
||||
extra_params: serde_json::json!({ "temperature": 0.7 }),
|
||||
})
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(result.models, vec!["model-a"]);
|
||||
let (method, uri, headers, body) = requests.recv().await.unwrap();
|
||||
assert_eq!(method, axum::http::Method::GET);
|
||||
assert_eq!(uri.path(), "/v1/models");
|
||||
assert!(body.is_empty());
|
||||
assert!(headers.get(axum::http::header::USER_AGENT).is_none());
|
||||
assert_eq!(headers.get("x-tenant").unwrap(), "tenant-a");
|
||||
assert_eq!(
|
||||
headers.get(axum::http::header::AUTHORIZATION).unwrap(),
|
||||
"Bearer secret"
|
||||
);
|
||||
server.abort();
|
||||
}
|
||||
}
|
||||
|
||||
@@ -459,6 +459,7 @@ pub fn dynamic_mcp(
|
||||
Error::Protocol(format!("MCP tool {} is missing input schema", wire.name))
|
||||
})?),
|
||||
};
|
||||
let parameters = normalize_mcp_parameters(&wire.name, parameters)?;
|
||||
let name = model_tool_name(&wire.name);
|
||||
let definition = ToolDefinition {
|
||||
name: name.clone(),
|
||||
@@ -477,6 +478,43 @@ pub fn dynamic_mcp(
|
||||
Ok(output)
|
||||
}
|
||||
|
||||
fn normalize_mcp_parameters(tool_name: &str, mut parameters: Value) -> Result<Value> {
|
||||
let schema = parameters
|
||||
.as_object_mut()
|
||||
.ok_or_else(|| invalid_mcp_parameters(tool_name))?;
|
||||
match schema.get("type") {
|
||||
Some(Value::String(schema_type)) if schema_type == "object" => return Ok(parameters),
|
||||
Some(_) => return Err(invalid_mcp_parameters(tool_name)),
|
||||
None => {}
|
||||
}
|
||||
let object_only_union = ["anyOf", "oneOf"].into_iter().any(|keyword| {
|
||||
schema
|
||||
.get(keyword)
|
||||
.and_then(Value::as_array)
|
||||
.is_some_and(|branches| {
|
||||
!branches.is_empty()
|
||||
&& branches.iter().all(|branch| {
|
||||
branch
|
||||
.as_object()
|
||||
.and_then(|branch| branch.get("type"))
|
||||
.and_then(Value::as_str)
|
||||
== Some("object")
|
||||
})
|
||||
})
|
||||
});
|
||||
if !object_only_union {
|
||||
return Err(invalid_mcp_parameters(tool_name));
|
||||
}
|
||||
schema.insert("type".into(), Value::String("object".into()));
|
||||
Ok(parameters)
|
||||
}
|
||||
|
||||
fn invalid_mcp_parameters(tool_name: &str) -> Error {
|
||||
Error::Protocol(format!(
|
||||
"MCP tool {tool_name} input schema must describe an object"
|
||||
))
|
||||
}
|
||||
|
||||
fn model_tool_name(name: &str) -> String {
|
||||
name.chars()
|
||||
.map(|character| {
|
||||
@@ -575,6 +613,115 @@ mod tests {
|
||||
.contains("duplicate MCP tool name after normalization: server_name-tool"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_normalizes_cursor_object_union_without_mutating_wire_schema() {
|
||||
let original_schema = serde_json::json!({
|
||||
"$schema": "https://json-schema.org/draft/2020-12/schema",
|
||||
"anyOf": [
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"rootPath": { "type": "string", "minLength": 1 }
|
||||
},
|
||||
"required": ["rootPath"],
|
||||
"additionalProperties": false
|
||||
},
|
||||
{
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"rootPaths": {
|
||||
"type": "array",
|
||||
"items": { "type": "string", "minLength": 1 },
|
||||
"minItems": 1
|
||||
}
|
||||
},
|
||||
"required": ["rootPaths"],
|
||||
"additionalProperties": false
|
||||
}
|
||||
]
|
||||
});
|
||||
let original_json = original_schema.to_string();
|
||||
let mut tool = direct_mcp_tool("cursor-app-control-move_agent_to_cloned_root");
|
||||
tool.input_schema_json = Some(original_json.clone());
|
||||
let request = pb::AgentRunRequest {
|
||||
mcp_tools: Some(pb::McpTools {
|
||||
mcp_tools: vec![tool],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let tools = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap();
|
||||
let (wire, definition) = tools
|
||||
.get("cursor-app-control-move_agent_to_cloned_root")
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(definition.parameters["type"], "object");
|
||||
assert_eq!(definition.parameters["anyOf"], original_schema["anyOf"]);
|
||||
assert_eq!(
|
||||
wire.input_schema_json.as_deref(),
|
||||
Some(original_json.as_str())
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_preserves_valid_object_schema() {
|
||||
let original_schema = serde_json::json!({
|
||||
"type": "object",
|
||||
"properties": {
|
||||
"query": { "type": "string" }
|
||||
},
|
||||
"required": ["query"],
|
||||
"additionalProperties": false
|
||||
});
|
||||
let mut tool = direct_mcp_tool("search");
|
||||
tool.input_schema_json = Some(original_schema.to_string());
|
||||
let request = pb::AgentRunRequest {
|
||||
mcp_tools: Some(pb::McpTools {
|
||||
mcp_tools: vec![tool],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let tools = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap();
|
||||
let (_, definition) = tools.get("search").unwrap();
|
||||
|
||||
assert_eq!(definition.parameters, original_schema);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn dynamic_mcp_rejects_schemas_that_are_not_provably_objects() {
|
||||
let invalid_schemas = [
|
||||
serde_json::Value::Null,
|
||||
serde_json::json!({ "type": "string" }),
|
||||
serde_json::json!({ "properties": { "query": { "type": "string" } } }),
|
||||
serde_json::json!({
|
||||
"anyOf": [
|
||||
{ "type": "object" },
|
||||
{ "type": "string" }
|
||||
]
|
||||
}),
|
||||
];
|
||||
|
||||
for schema in invalid_schemas {
|
||||
let mut tool = direct_mcp_tool("unsafe_schema");
|
||||
tool.input_schema_json = Some(schema.to_string());
|
||||
let request = pb::AgentRunRequest {
|
||||
mcp_tools: Some(pb::McpTools {
|
||||
mcp_tools: vec![tool],
|
||||
}),
|
||||
..Default::default()
|
||||
};
|
||||
|
||||
let error = dynamic_mcp(&request, &pb::RequestContext::default()).unwrap_err();
|
||||
assert!(
|
||||
error
|
||||
.to_string()
|
||||
.contains("MCP tool unsafe_schema input schema must describe an object"),
|
||||
"unexpected error for {schema}: {error}"
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn meta_mcp_routes_projects_descriptor_routing_without_runtime_discovery() {
|
||||
let context = pb::RequestContext {
|
||||
|
||||
@@ -129,11 +129,19 @@ impl CursorSession {
|
||||
checkpoint_worker_open = false;
|
||||
}
|
||||
Input::Completion(completion) => {
|
||||
self.forward_completion(completion, &mut completions)
|
||||
.await?;
|
||||
if let Some(completion) = self
|
||||
.forward_completion(completion, &mut completions)
|
||||
.await?
|
||||
{
|
||||
ready.push_back(completion);
|
||||
}
|
||||
}
|
||||
Input::CompletionResult(Some(result)) => {
|
||||
self.forward_completion(result?, &mut completions).await?;
|
||||
if let Some(completion) =
|
||||
self.forward_completion(result?, &mut completions).await?
|
||||
{
|
||||
ready.push_back(completion);
|
||||
}
|
||||
}
|
||||
Input::CompletionResult(None) => {
|
||||
return Err(Error::Protocol("tool result channel closed".into()));
|
||||
@@ -578,7 +586,7 @@ impl CursorSession {
|
||||
&self,
|
||||
mut completion: ToolCompletion,
|
||||
completions: &mut HashMap<String, ToolCompletion>,
|
||||
) -> Result<()> {
|
||||
) -> Result<Option<ToolCompletion>> {
|
||||
if let Some(image) = completion.take_read_image() {
|
||||
let blob_id = self.store.put_blob(&image.data, &[]).await?;
|
||||
completion.persist_read_image(&blob_id, &image)?;
|
||||
@@ -600,7 +608,14 @@ impl CursorSession {
|
||||
.commands
|
||||
.send(ClientCommand::ToolResult(result.clone()))
|
||||
.await
|
||||
.map_err(|_| Error::RunNotFound(self.context.request_id.clone()))
|
||||
.map_err(|_| Error::RunNotFound(self.context.request_id.clone()))?;
|
||||
let Some(dispatched) = self.tools.continue_after(&result.call_id).await? else {
|
||||
return Ok(None);
|
||||
};
|
||||
for message in dispatched.messages {
|
||||
self.handle.emit(&message)?;
|
||||
}
|
||||
Ok(dispatched.completion)
|
||||
}
|
||||
|
||||
async fn forward_injection(&mut self, action: pb::InjectContextAction) -> Result<()> {
|
||||
|
||||
@@ -20,6 +20,13 @@ pub(crate) fn path(call: &ToolCall) -> Result<String> {
|
||||
string(call, field)
|
||||
}
|
||||
|
||||
pub(crate) fn execution_path(call: &ToolCall) -> Result<Option<String>> {
|
||||
match normalized(&call.name).as_str() {
|
||||
"write" | "strreplace" | "editnotebook" => path(call).map(Some),
|
||||
_ => Ok(None),
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) fn after_read(
|
||||
call: &ToolCall,
|
||||
result: &pb::ReadResult,
|
||||
|
||||
@@ -1,11 +1,19 @@
|
||||
use std::collections::{BTreeMap, HashSet};
|
||||
use std::{
|
||||
collections::{BTreeMap, HashSet},
|
||||
sync::Arc,
|
||||
};
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
|
||||
pub mod codec;
|
||||
mod dispatch;
|
||||
pub(crate) mod edit;
|
||||
pub(crate) mod result;
|
||||
pub mod runtime;
|
||||
mod schedule;
|
||||
pub(crate) mod stream;
|
||||
#[cfg(test)]
|
||||
mod tests;
|
||||
|
||||
use crate::{
|
||||
model::{CanonicalMessage, MessageContent, Role, ToolCall},
|
||||
@@ -14,6 +22,7 @@ use crate::{
|
||||
};
|
||||
|
||||
use self::result::{ToolCompletion, ToolResultSender};
|
||||
use self::schedule::{DeferredEdit, EditSchedule};
|
||||
use super::{interaction, proto::agent::v1 as pb};
|
||||
use runtime::{CursorToolRuntime, ExecContext};
|
||||
|
||||
@@ -23,6 +32,7 @@ pub struct ToolDispatcher {
|
||||
results: ToolResultSender,
|
||||
search: WebSearch,
|
||||
fetch: WebFetch,
|
||||
edit_schedule: Arc<Mutex<EditSchedule>>,
|
||||
}
|
||||
|
||||
pub struct DispatchedTool {
|
||||
@@ -54,6 +64,7 @@ impl ToolDispatcher {
|
||||
results,
|
||||
search: WebSearch::built_in(),
|
||||
fetch: WebFetch::built_in(),
|
||||
edit_schedule: Arc::new(Mutex::new(EditSchedule::default())),
|
||||
}
|
||||
}
|
||||
|
||||
@@ -74,20 +85,62 @@ impl ToolDispatcher {
|
||||
if state.completed.contains(&call.call_id) {
|
||||
continue;
|
||||
}
|
||||
let message_index = first_tool_index + position;
|
||||
let publish_started = !state.started.contains(&call.call_id);
|
||||
let edit_path = if dynamic_mcp.contains_key(&call.name) {
|
||||
None
|
||||
} else {
|
||||
edit::execution_path(call)?
|
||||
};
|
||||
if let Some(path) = edit_path {
|
||||
let next = self.edit_schedule.lock().await.start_or_defer(
|
||||
path,
|
||||
DeferredEdit {
|
||||
call: call.clone(),
|
||||
message_index,
|
||||
publish_started,
|
||||
context: context.clone(),
|
||||
},
|
||||
);
|
||||
let Some(next) = next else {
|
||||
continue;
|
||||
};
|
||||
dispatched.push(
|
||||
self.start(
|
||||
&next.call,
|
||||
next.message_index,
|
||||
next.publish_started,
|
||||
dynamic_mcp,
|
||||
&next.context,
|
||||
)
|
||||
.await?,
|
||||
);
|
||||
continue;
|
||||
}
|
||||
dispatched.push(
|
||||
self.start(
|
||||
call,
|
||||
first_tool_index + position,
|
||||
!state.started.contains(&call.call_id),
|
||||
dynamic_mcp,
|
||||
context,
|
||||
)
|
||||
.await?,
|
||||
self.start(call, message_index, publish_started, dynamic_mcp, context)
|
||||
.await?,
|
||||
);
|
||||
}
|
||||
Ok(dispatched)
|
||||
}
|
||||
|
||||
pub(crate) async fn continue_after(&self, call_id: &str) -> Result<Option<DispatchedTool>> {
|
||||
let next = self.edit_schedule.lock().await.complete(call_id)?;
|
||||
let Some(next) = next else {
|
||||
return Ok(None);
|
||||
};
|
||||
self.start(
|
||||
&next.call,
|
||||
next.message_index,
|
||||
next.publish_started,
|
||||
&BTreeMap::new(),
|
||||
&next.context,
|
||||
)
|
||||
.await
|
||||
.map(Some)
|
||||
}
|
||||
|
||||
async fn start(
|
||||
&self,
|
||||
call: &ToolCall,
|
||||
|
||||
@@ -0,0 +1,68 @@
|
||||
use std::collections::{HashMap, VecDeque};
|
||||
|
||||
use crate::{model::ToolCall, Error, Result};
|
||||
|
||||
use super::runtime::ExecContext;
|
||||
|
||||
#[derive(Default)]
|
||||
pub(super) struct EditSchedule {
|
||||
paths: HashMap<String, EditPathQueue>,
|
||||
active_paths: HashMap<String, String>,
|
||||
}
|
||||
|
||||
struct EditPathQueue {
|
||||
active_call_id: String,
|
||||
waiting: VecDeque<DeferredEdit>,
|
||||
}
|
||||
|
||||
pub(super) struct DeferredEdit {
|
||||
pub call: ToolCall,
|
||||
pub message_index: usize,
|
||||
pub publish_started: bool,
|
||||
pub context: ExecContext,
|
||||
}
|
||||
|
||||
impl EditSchedule {
|
||||
pub fn start_or_defer(&mut self, path: String, edit: DeferredEdit) -> Option<DeferredEdit> {
|
||||
if let Some(queue) = self.paths.get_mut(&path) {
|
||||
queue.waiting.push_back(edit);
|
||||
return None;
|
||||
}
|
||||
self.active_paths
|
||||
.insert(edit.call.call_id.clone(), path.clone());
|
||||
self.paths.insert(
|
||||
path,
|
||||
EditPathQueue {
|
||||
active_call_id: edit.call.call_id.clone(),
|
||||
waiting: VecDeque::new(),
|
||||
},
|
||||
);
|
||||
Some(edit)
|
||||
}
|
||||
|
||||
pub fn complete(&mut self, call_id: &str) -> Result<Option<DeferredEdit>> {
|
||||
let Some(path) = self.active_paths.remove(call_id) else {
|
||||
return Ok(None);
|
||||
};
|
||||
let queue = self.paths.get_mut(&path).ok_or_else(|| {
|
||||
Error::Protocol(format!("active edit path disappeared for call {call_id}"))
|
||||
})?;
|
||||
if queue.active_call_id != call_id {
|
||||
return Err(Error::Protocol(format!(
|
||||
"edit path is active for {}, not {call_id}",
|
||||
queue.active_call_id
|
||||
)));
|
||||
}
|
||||
match queue.waiting.pop_front() {
|
||||
Some(next) => {
|
||||
queue.active_call_id = next.call.call_id.clone();
|
||||
self.active_paths.insert(next.call.call_id.clone(), path);
|
||||
Ok(Some(next))
|
||||
}
|
||||
None => {
|
||||
self.paths.remove(&path);
|
||||
Ok(None)
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,139 @@
|
||||
use super::*;
|
||||
use serde_json::json;
|
||||
|
||||
fn edit_call(index: usize, call_id: &str, path: &str, old: &str, new: &str) -> ToolCall {
|
||||
ToolCall {
|
||||
index,
|
||||
call_id: call_id.into(),
|
||||
model_call_id: "model:0".into(),
|
||||
name: "StrReplace".into(),
|
||||
arguments_text: String::new(),
|
||||
arguments: json!({
|
||||
"path": path,
|
||||
"old_string": old,
|
||||
"new_string": new,
|
||||
}),
|
||||
}
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn same_path_edits_start_one_at_a_time() {
|
||||
let runtime = CursorToolRuntime::default();
|
||||
let dispatcher = ToolDispatcher::new(runtime.clone());
|
||||
let calls = [
|
||||
edit_call(0, "first", "/tmp/a.txt", "left", "LEFT"),
|
||||
edit_call(1, "second", "/tmp/a.txt", "right", "RIGHT"),
|
||||
edit_call(2, "other", "/tmp/b.txt", "other", "OTHER"),
|
||||
];
|
||||
|
||||
let dispatched = dispatcher
|
||||
.start_batch(
|
||||
&calls,
|
||||
ToolBatchState {
|
||||
completed: &HashSet::new(),
|
||||
started: &HashSet::new(),
|
||||
response_text: "",
|
||||
response_thinking: "",
|
||||
},
|
||||
&[],
|
||||
&BTreeMap::new(),
|
||||
&ExecContext::default(),
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
assert_eq!(dispatched.len(), 2);
|
||||
assert_eq!(exec(&dispatched[0]).exec_id, "first");
|
||||
assert_eq!(exec(&dispatched[1]).exec_id, "other");
|
||||
|
||||
let mut file = "left right\n".to_string();
|
||||
let first_write = advance_read(&runtime, exec(&dispatched[0]).id, &file).await;
|
||||
file = write_text(&first_write);
|
||||
assert_eq!(file, "LEFT right\n");
|
||||
complete_write(&runtime, &first_write).await;
|
||||
|
||||
let second = dispatcher
|
||||
.continue_after("first")
|
||||
.await
|
||||
.unwrap()
|
||||
.expect("second same-path edit should start after the first completes");
|
||||
assert_eq!(exec(&second).exec_id, "second");
|
||||
let second_write = advance_read(&runtime, exec(&second).id, &file).await;
|
||||
file = write_text(&second_write);
|
||||
assert_eq!(file, "LEFT RIGHT\n");
|
||||
complete_write(&runtime, &second_write).await;
|
||||
assert!(dispatcher.continue_after("second").await.unwrap().is_none());
|
||||
}
|
||||
|
||||
fn exec(dispatched: &DispatchedTool) -> &pb::ExecServerMessage {
|
||||
dispatched
|
||||
.messages
|
||||
.iter()
|
||||
.find_map(|message| match message.message.as_ref() {
|
||||
Some(pb::agent_server_message::Message::ExecServerMessage(exec)) => Some(exec),
|
||||
_ => None,
|
||||
})
|
||||
.expect("dispatched edit should contain an Exec request")
|
||||
}
|
||||
|
||||
async fn advance_read(
|
||||
runtime: &CursorToolRuntime,
|
||||
id: u32,
|
||||
content: &str,
|
||||
) -> pb::ExecServerMessage {
|
||||
let event = codec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id,
|
||||
message: Some(pb::exec_client_message::Message::ReadResult(
|
||||
pb::ReadResult {
|
||||
result: Some(pb::read_result::Result::Success(pb::ReadSuccess {
|
||||
output: Some(pb::read_success::Output::Content(content.into())),
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
runtime,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
let codec::ClientExecEvent::Message(message) = event else {
|
||||
panic!("edit read should advance to a write")
|
||||
};
|
||||
let Some(pb::agent_server_message::Message::ExecServerMessage(exec)) = message.message else {
|
||||
panic!("edit read should emit an Exec write request")
|
||||
};
|
||||
exec
|
||||
}
|
||||
|
||||
fn write_text(exec: &pb::ExecServerMessage) -> String {
|
||||
let Some(pb::exec_server_message::Message::WriteArgs(args)) = exec.message.as_ref() else {
|
||||
panic!("expected WriteArgs")
|
||||
};
|
||||
args.file_text.clone()
|
||||
}
|
||||
|
||||
async fn complete_write(runtime: &CursorToolRuntime, exec: &pb::ExecServerMessage) {
|
||||
let Some(pb::exec_server_message::Message::WriteArgs(args)) = exec.message.as_ref() else {
|
||||
panic!("expected WriteArgs")
|
||||
};
|
||||
let event = codec::client_event(
|
||||
&pb::ExecClientMessage {
|
||||
id: exec.id,
|
||||
message: Some(pb::exec_client_message::Message::WriteResult(
|
||||
pb::WriteResult {
|
||||
result: Some(pb::write_result::Result::Success(pb::WriteSuccess {
|
||||
path: args.path.clone(),
|
||||
..Default::default()
|
||||
})),
|
||||
},
|
||||
)),
|
||||
..Default::default()
|
||||
},
|
||||
runtime,
|
||||
)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(matches!(event, codec::ClientExecEvent::Completed(_)));
|
||||
}
|
||||
@@ -51,6 +51,7 @@ pub struct ProviderEndpoint {
|
||||
pub name: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub base_url: String,
|
||||
pub api_key: Option<String>,
|
||||
pub has_api_key: bool,
|
||||
pub custom_headers: serde_json::Value,
|
||||
pub extra_params: serde_json::Value,
|
||||
@@ -61,7 +62,6 @@ pub struct ProviderEndpoint {
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ProviderEndpointSecret {
|
||||
pub endpoint: ProviderEndpoint,
|
||||
pub api_key: String,
|
||||
pub custom_headers: serde_json::Value,
|
||||
}
|
||||
|
||||
|
||||
@@ -82,7 +82,7 @@ impl Provider for ProviderRouter {
|
||||
ProviderType::Anthropic => ProviderKind::Anthropic,
|
||||
},
|
||||
request_url,
|
||||
api_key: endpoint.api_key,
|
||||
api_key: endpoint.endpoint.api_key.clone().unwrap_or_default(),
|
||||
custom_headers: custom_headers(&endpoint.custom_headers)?,
|
||||
max_output_tokens: model.max_output_tokens,
|
||||
request_timeout,
|
||||
|
||||
@@ -127,10 +127,15 @@ impl Store {
|
||||
.provider(provider_id)
|
||||
.await?
|
||||
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
|
||||
let api_key = input.api_key.as_deref().unwrap_or(¤t.api_key);
|
||||
let api_key = input
|
||||
.api_key
|
||||
.as_deref()
|
||||
.or(current.endpoint.api_key.as_deref())
|
||||
.unwrap_or_default();
|
||||
let custom_headers = merge_custom_headers(¤t.custom_headers, &input.custom_headers)?;
|
||||
let base_url = normalize_base_url(&input.base_url)?;
|
||||
let identity_changed = base_url != current.endpoint.base_url || api_key != current.api_key;
|
||||
let identity_changed = base_url != current.endpoint.base_url
|
||||
|| api_key != current.endpoint.api_key.as_deref().unwrap_or_default();
|
||||
let models = if identity_changed {
|
||||
sqlx::query("SELECT * FROM provider_models WHERE provider_id = ?")
|
||||
.bind(provider_id)
|
||||
@@ -276,7 +281,7 @@ impl Store {
|
||||
for input in inputs {
|
||||
let hash = model_hash(
|
||||
&provider.endpoint.base_url,
|
||||
&provider.api_key,
|
||||
provider.endpoint.api_key.as_deref().unwrap_or_default(),
|
||||
input.endpoint_type,
|
||||
&input.model_id,
|
||||
)?;
|
||||
@@ -333,7 +338,7 @@ impl Store {
|
||||
.expect("model provider must exist");
|
||||
let next_hash = model_hash(
|
||||
&provider.endpoint.base_url,
|
||||
&provider.api_key,
|
||||
provider.endpoint.api_key.as_deref().unwrap_or_default(),
|
||||
input.endpoint_type,
|
||||
&input.model_id,
|
||||
)?;
|
||||
@@ -485,6 +490,7 @@ fn validate_model_batch(inputs: &[ProviderModelInput]) -> Result<()> {
|
||||
|
||||
fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
|
||||
let api_key: String = row.try_get("api_key")?;
|
||||
let has_api_key = !api_key.is_empty();
|
||||
let headers: serde_json::Value = serde_json::from_str(row.try_get("custom_headers_json")?)?;
|
||||
let extra_params: serde_json::Value = serde_json::from_str(row.try_get("extra_params_json")?)?;
|
||||
Ok(ProviderEndpoint {
|
||||
@@ -492,7 +498,8 @@ fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
|
||||
name: row.try_get("name")?,
|
||||
provider_type: ProviderType::from_str(row.try_get("provider_type")?)?,
|
||||
base_url: row.try_get("base_url")?,
|
||||
has_api_key: !api_key.is_empty(),
|
||||
api_key: has_api_key.then_some(api_key),
|
||||
has_api_key,
|
||||
custom_headers: redact_custom_headers(&headers),
|
||||
extra_params,
|
||||
created_at_ms: row.try_get("created_at_ms")?,
|
||||
@@ -501,12 +508,10 @@ fn endpoint_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpoint> {
|
||||
}
|
||||
|
||||
fn secret_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ProviderEndpointSecret> {
|
||||
let api_key: String = row.try_get("api_key")?;
|
||||
let custom_headers: serde_json::Value =
|
||||
serde_json::from_str(row.try_get("custom_headers_json")?)?;
|
||||
Ok(ProviderEndpointSecret {
|
||||
endpoint: endpoint_from_row(row)?,
|
||||
api_key,
|
||||
custom_headers,
|
||||
})
|
||||
}
|
||||
@@ -847,6 +852,69 @@ mod tests {
|
||||
assert_eq!(detached, None);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn updating_provider_without_changing_api_key_preserves_model_hashes() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("provider-key-keep.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let (created_provider, original) = store
|
||||
.create_provider_with_model(&provider(), &model("model-a"))
|
||||
.await
|
||||
.unwrap();
|
||||
|
||||
// Editor keeps the configured key: sending it back must not rehash models.
|
||||
store
|
||||
.update_provider(created_provider.provider_id, &provider())
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(store
|
||||
.provider_model(&original.model_hash)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some());
|
||||
|
||||
// Editor cleared the field: keep the current key, still no rehash.
|
||||
let mut without_key = provider();
|
||||
without_key.api_key = None;
|
||||
store
|
||||
.update_provider(created_provider.provider_id, &without_key)
|
||||
.await
|
||||
.unwrap();
|
||||
assert!(store
|
||||
.provider_model(&original.model_hash)
|
||||
.await
|
||||
.unwrap()
|
||||
.is_some());
|
||||
assert_eq!(store.provider_models(false).await.unwrap().len(), 1);
|
||||
}
|
||||
|
||||
#[tokio::test]
|
||||
async fn providers_expose_the_configured_api_key_for_editing() {
|
||||
let directory = tempfile::tempdir().unwrap();
|
||||
let store = Store::connect(&format!(
|
||||
"sqlite://{}",
|
||||
directory.path().join("provider-key-echo.db").display()
|
||||
))
|
||||
.await
|
||||
.unwrap();
|
||||
let created = store.create_provider(&provider()).await.unwrap();
|
||||
assert_eq!(created.api_key.as_deref(), Some("secret"));
|
||||
let listed = store.providers().await.unwrap();
|
||||
assert_eq!(listed.len(), 1);
|
||||
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 empty = store.create_provider(&without_key).await.unwrap();
|
||||
assert_eq!(empty.api_key, None);
|
||||
assert!(!empty.has_api_key);
|
||||
assert_eq!(store.providers().await.unwrap().len(), 2);
|
||||
}
|
||||
|
||||
async fn insert_call(store: &Store, provider: &ProviderEndpoint, model: &ProviderModel) {
|
||||
sqlx::query(
|
||||
"INSERT INTO llm_calls(call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type, provider_url, request_type, request_url, model_id, display_name, status, created_at_ms, message_count, tool_count, detailed) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)",
|
||||
|
||||
Reference in New Issue
Block a user