mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
refactor: align project directory architecture
This commit is contained in:
+1
-1
@@ -2,7 +2,7 @@ use std::{env, path::PathBuf};
|
||||
|
||||
fn main() {
|
||||
let manifest = PathBuf::from(env::var("CARGO_MANIFEST_DIR").expect("manifest directory"));
|
||||
let proto_dir = manifest.join("../scripts/cursor-proto/proto");
|
||||
let proto_dir = manifest.join("../protocols/cursor");
|
||||
let protos = [proto_dir.join("agent_v1.proto")];
|
||||
let aiserver_proto = proto_dir.join("aiserver_v1.proto");
|
||||
|
||||
|
||||
@@ -1,19 +0,0 @@
|
||||
use tokio::sync::oneshot;
|
||||
|
||||
use crate::model::{CanonicalMessage, RuntimeEvent, ToolResult};
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct MessageInsertion {
|
||||
pub messages: Vec<CanonicalMessage>,
|
||||
pub delivered: oneshot::Sender<()>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum ClientCommand {
|
||||
ToolResult(ToolResult),
|
||||
InterruptWithMessage(CanonicalMessage),
|
||||
RuntimeEvent(RuntimeEvent),
|
||||
InsertMessages(MessageInsertion),
|
||||
ClientClosed { error: String },
|
||||
Cancel,
|
||||
}
|
||||
@@ -1,7 +0,0 @@
|
||||
mod command;
|
||||
mod event;
|
||||
mod session;
|
||||
|
||||
pub use command::*;
|
||||
pub use event::*;
|
||||
pub use session::*;
|
||||
@@ -1,28 +0,0 @@
|
||||
use tokio::sync::mpsc;
|
||||
|
||||
use super::{ClientCommand, ClientEvent};
|
||||
|
||||
pub struct ClientPort {
|
||||
pub commands: mpsc::Receiver<ClientCommand>,
|
||||
pub events: mpsc::Sender<ClientEvent>,
|
||||
}
|
||||
|
||||
pub struct ClientSession {
|
||||
pub commands: mpsc::Sender<ClientCommand>,
|
||||
pub events: mpsc::Receiver<ClientEvent>,
|
||||
}
|
||||
|
||||
pub fn session(capacity: usize) -> (ClientPort, ClientSession) {
|
||||
let (commands_tx, commands_rx) = mpsc::channel(capacity);
|
||||
let (events_tx, events_rx) = mpsc::channel(capacity);
|
||||
(
|
||||
ClientPort {
|
||||
commands: commands_rx,
|
||||
events: events_tx,
|
||||
},
|
||||
ClientSession {
|
||||
commands: commands_tx,
|
||||
events: events_rx,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -193,7 +193,7 @@ impl CursorActor {
|
||||
return;
|
||||
}
|
||||
let cancellation = handle.cancellation();
|
||||
let (port, core) = crate::client::session(256);
|
||||
let (port, core) = crate::run::session(256);
|
||||
let core_commands = core.commands.clone();
|
||||
let actor = RunActor::new(
|
||||
dependencies.store.clone(),
|
||||
|
||||
@@ -3,7 +3,6 @@ use std::collections::{BTreeMap, HashMap, HashSet, VecDeque};
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::{
|
||||
client::{ClientCommand, ClientEvent, ClientSession, CommitCause},
|
||||
cursor::{
|
||||
blob_sync::BlobSynchronizer,
|
||||
checkpoint::{
|
||||
@@ -26,7 +25,7 @@ use crate::{
|
||||
},
|
||||
},
|
||||
model::{ToolCall, ToolRoundId, Usage},
|
||||
run::{RunFailure, RunOutcome},
|
||||
run::{ClientCommand, ClientEvent, ClientSession, CommitCause, RunFailure, RunOutcome},
|
||||
store::{Store, ToolRoundStatus},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -3,7 +3,7 @@
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
model::ToolCall,
|
||||
web::{WebFetch, WebSearch},
|
||||
search::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -117,7 +117,7 @@ mod tests {
|
||||
use crate::{
|
||||
cursor::{proto::agent::v1 as pb, tools::result::tool_result_channel},
|
||||
model::ToolCall,
|
||||
web::{HtmlEngine, WebFetch, WebSearch},
|
||||
search::{HtmlEngine, WebFetch, WebSearch},
|
||||
};
|
||||
|
||||
use super::{resume, InteractionContinuation, PendingInteraction};
|
||||
|
||||
@@ -9,8 +9,8 @@ use std::collections::BTreeMap;
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::ToolCall,
|
||||
search::{WebFetch, WebSearch},
|
||||
store::Store,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
|
||||
@@ -1,79 +1,30 @@
|
||||
//! Asynchronous dispatch for the application-owned Semble search tools.
|
||||
//! Cursor tool orchestration for application-owned Semble search.
|
||||
|
||||
use std::sync::Arc;
|
||||
|
||||
use semble_core::{ContentType, FindRelatedRequest, SearchEngine, SearchRequest, SembleConfig};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
use crate::{model::ToolCall, store::Store, Error, Result};
|
||||
|
||||
use super::ToolStart;
|
||||
use crate::cursor::tools::{
|
||||
result::{self, ToolResultSender},
|
||||
runtime::now_ms,
|
||||
use crate::{
|
||||
cursor::tools::{
|
||||
result::{self, ToolResultSender},
|
||||
runtime::now_ms,
|
||||
},
|
||||
model::ToolCall,
|
||||
search,
|
||||
store::Store,
|
||||
Result,
|
||||
};
|
||||
|
||||
static ENGINE: OnceCell<Arc<SearchEngine>> = OnceCell::const_new();
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
enum ContentSelection {
|
||||
#[default]
|
||||
Code,
|
||||
Docs,
|
||||
Config,
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct SearchArguments {
|
||||
query: String,
|
||||
repo: String,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct FindRelatedArguments {
|
||||
repo: String,
|
||||
file_path: String,
|
||||
line: usize,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
use super::ToolStart;
|
||||
|
||||
pub(super) fn start(
|
||||
results: &ToolResultSender,
|
||||
call: &ToolCall,
|
||||
store: Option<Store>,
|
||||
) -> Result<ToolStart> {
|
||||
let operation = match super::normalized(&call.name).as_str() {
|
||||
"semblesearch" => Operation::Search(serde_json::from_value(call.arguments.clone())?),
|
||||
"semblefindrelated" => {
|
||||
Operation::FindRelated(serde_json::from_value(call.arguments.clone())?)
|
||||
}
|
||||
_ => {
|
||||
return Err(Error::Protocol(format!(
|
||||
"unsupported Semble tool: {}",
|
||||
call.name
|
||||
)))
|
||||
}
|
||||
};
|
||||
let tool_name = super::normalized(&call.name);
|
||||
let arguments = call.arguments.clone();
|
||||
let call = call.clone();
|
||||
let results = results.clone();
|
||||
let started_at_ms = now_ms();
|
||||
tokio::spawn(async move {
|
||||
let output = execute(operation, store).await;
|
||||
let output = search::execute_semble(&tool_name, arguments, store).await;
|
||||
match result::semble(&call, started_at_ms, output) {
|
||||
Ok(completion) => results.send(completion),
|
||||
Err(error) => results.send_error(error),
|
||||
@@ -84,117 +35,3 @@ pub(super) fn start(
|
||||
completion: None,
|
||||
})
|
||||
}
|
||||
|
||||
enum Operation {
|
||||
Search(SearchArguments),
|
||||
FindRelated(FindRelatedArguments),
|
||||
}
|
||||
|
||||
async fn execute(operation: Operation, store: Option<Store>) -> std::result::Result<Value, String> {
|
||||
let engine = engine(store).await.map_err(|error| error.to_string())?;
|
||||
tokio::task::spawn_blocking(move || match operation {
|
||||
Operation::Search(arguments) => engine
|
||||
.search(SearchRequest {
|
||||
query: arguments.query,
|
||||
repo: arguments.repo.into(),
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
Operation::FindRelated(arguments) => engine
|
||||
.find_related(FindRelatedRequest {
|
||||
repo: arguments.repo.into(),
|
||||
file_path: arguments.file_path,
|
||||
line: arguments.line,
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
})
|
||||
.await
|
||||
.map_err(|error| format!("Semble search worker failed: {error}"))?
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
async fn engine(store: Option<Store>) -> Result<Arc<SearchEngine>> {
|
||||
ENGINE
|
||||
.get_or_try_init(|| async move {
|
||||
let builder = match store {
|
||||
Some(store) => crate::network::blocking_client_builder(&store).await?,
|
||||
None => reqwest::blocking::Client::builder().use_native_tls(),
|
||||
};
|
||||
tokio::task::spawn_blocking(move || {
|
||||
let client = builder.build()?;
|
||||
SearchEngine::load_default_with_client(SembleConfig::default(), &client)
|
||||
.map(Arc::new)
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))?
|
||||
})
|
||||
.await
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn json_value(response: semble_core::SearchResponse) -> semble_core::Result<serde_json::Value> {
|
||||
serde_json::to_value(response)
|
||||
.map_err(|error| semble_core::Error::Serialization(error.to_string()))
|
||||
}
|
||||
|
||||
fn content(selection: ContentSelection) -> Vec<ContentType> {
|
||||
match selection {
|
||||
ContentSelection::Code => vec![ContentType::Code],
|
||||
ContentSelection::Docs => vec![ContentType::Docs],
|
||||
ContentSelection::Config => vec![ContentType::Config],
|
||||
ContentSelection::All => vec![ContentType::Code, ContentType::Docs, ContentType::Config],
|
||||
}
|
||||
}
|
||||
|
||||
fn default_top_k() -> usize {
|
||||
5
|
||||
}
|
||||
|
||||
fn default_snippet_lines() -> Option<usize> {
|
||||
Some(10)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn search_arguments_use_code_search_defaults() {
|
||||
let arguments: SearchArguments = serde_json::from_value(json!({
|
||||
"query": "request persistence",
|
||||
"repo": "/tmp/repo"
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.top_k, 5);
|
||||
assert_eq!(arguments.max_snippet_lines, Some(10));
|
||||
assert!(matches!(arguments.content, ContentSelection::Code));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn find_related_does_not_require_a_ui_description() {
|
||||
let arguments: FindRelatedArguments = serde_json::from_value(json!({
|
||||
"repo": "/tmp/repo",
|
||||
"file_path": "src/auth.ts",
|
||||
"line": 42
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.file_path, "src/auth.ts");
|
||||
assert_eq!(arguments.line, 42);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn all_content_expands_to_every_indexed_scope() {
|
||||
assert_eq!(
|
||||
content(ContentSelection::All),
|
||||
vec![ContentType::Code, ContentType::Docs, ContentType::Config]
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
@@ -18,8 +18,8 @@ mod tests;
|
||||
|
||||
use crate::{
|
||||
model::{CanonicalMessage, MessageContent, Role, ToolCall},
|
||||
search::{WebFetch, WebSearch},
|
||||
store::Store,
|
||||
web::{WebFetch, WebSearch},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
|
||||
@@ -1,6 +1,6 @@
|
||||
use crate::{
|
||||
cursor::{interaction, proto::agent::v1 as pb},
|
||||
web::{FetchedPage, SearchHit},
|
||||
search::{FetchedPage, SearchHit},
|
||||
Error, Result,
|
||||
};
|
||||
|
||||
@@ -335,7 +335,7 @@ mod tests {
|
||||
use crate::{
|
||||
cursor::proto::agent::v1 as pb,
|
||||
model::ToolCall,
|
||||
web::{FetchedPage, SearchHit},
|
||||
search::{FetchedPage, SearchHit},
|
||||
};
|
||||
|
||||
use super::{complete_web_fetch, complete_web_search, PendingInteraction};
|
||||
|
||||
+1
-2
@@ -1,5 +1,4 @@
|
||||
pub mod app;
|
||||
pub mod client;
|
||||
pub mod config;
|
||||
pub mod control;
|
||||
pub mod cursor;
|
||||
@@ -9,8 +8,8 @@ pub mod model;
|
||||
pub mod network;
|
||||
pub mod provider;
|
||||
pub mod run;
|
||||
pub mod search;
|
||||
pub mod store;
|
||||
pub mod web;
|
||||
|
||||
pub use app::App;
|
||||
pub use config::Config;
|
||||
|
||||
@@ -1,28 +0,0 @@
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct CursorRunTraceSummary {
|
||||
pub request_id: String,
|
||||
pub conversation_id: Option<String>,
|
||||
pub route: String,
|
||||
pub model_id: Option<String>,
|
||||
pub status: String,
|
||||
pub request_bytes: i64,
|
||||
pub response_bytes: i64,
|
||||
pub response_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub received_at_ms: i64,
|
||||
pub first_response_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub error_message: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct CursorRunTraceArtifact {
|
||||
pub seq: i64,
|
||||
pub artifact_type: String,
|
||||
pub source: String,
|
||||
pub metadata: serde_json::Value,
|
||||
pub created_at_ms: i64,
|
||||
pub data: Vec<u8>,
|
||||
}
|
||||
@@ -1,93 +0,0 @@
|
||||
use serde::Serialize;
|
||||
|
||||
use super::{ProviderType, Usage};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct NewLlmCall {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub provider_url: String,
|
||||
pub request_type: ProviderType,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: bool,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) struct LlmCallUsageAnchor {
|
||||
pub request_type: ProviderType,
|
||||
pub usage: Usage,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallSummary {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: Option<String>,
|
||||
pub provider_type: String,
|
||||
pub provider_url: String,
|
||||
pub request_type: String,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: Option<bool>,
|
||||
pub status: String,
|
||||
pub finish_reason: Option<String>,
|
||||
pub created_at_ms: i64,
|
||||
pub request_started_at_ms: Option<i64>,
|
||||
pub response_headers_at_ms: Option<i64>,
|
||||
pub first_event_at_ms: Option<i64>,
|
||||
pub first_text_at_ms: Option<i64>,
|
||||
pub first_valid_response_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub queue_ms: Option<i64>,
|
||||
pub ttfb_ms: Option<i64>,
|
||||
pub ttft_ms: Option<i64>,
|
||||
pub ttfr_ms: Option<i64>,
|
||||
pub duration_ms: Option<i64>,
|
||||
pub input_tokens: Option<i64>,
|
||||
pub output_tokens: Option<i64>,
|
||||
pub total_tokens: Option<i64>,
|
||||
pub cache_read_tokens: Option<i64>,
|
||||
pub cache_write_tokens: Option<i64>,
|
||||
pub reasoning_tokens: Option<i64>,
|
||||
pub usage: Option<serde_json::Value>,
|
||||
pub message_count: i64,
|
||||
pub tool_count: i64,
|
||||
pub request_bytes: Option<i64>,
|
||||
pub response_bytes: i64,
|
||||
pub stream_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub error_kind: Option<String>,
|
||||
pub error_message: Option<String>,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallRequest {
|
||||
pub headers: serde_json::Value,
|
||||
pub body: serde_json::Value,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallResponseChunk {
|
||||
pub seq: i64,
|
||||
pub received_offset_ms: i64,
|
||||
pub data: String,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
@@ -1,31 +1,25 @@
|
||||
mod configuration;
|
||||
mod conversation;
|
||||
mod cursor_trace;
|
||||
mod inference;
|
||||
mod llm_call;
|
||||
mod message;
|
||||
mod model_spec;
|
||||
mod overview;
|
||||
mod observability;
|
||||
mod projection;
|
||||
mod run;
|
||||
mod runtime_tag;
|
||||
mod token_count;
|
||||
mod tool;
|
||||
mod tool_result_replay;
|
||||
mod usage;
|
||||
|
||||
pub use configuration::*;
|
||||
pub use conversation::*;
|
||||
pub use cursor_trace::*;
|
||||
pub use inference::*;
|
||||
pub use llm_call::*;
|
||||
pub use message::*;
|
||||
pub use model_spec::*;
|
||||
pub use overview::*;
|
||||
pub use observability::*;
|
||||
pub use projection::*;
|
||||
pub use run::*;
|
||||
pub use runtime_tag::*;
|
||||
pub(crate) use token_count::*;
|
||||
pub use tool::*;
|
||||
pub(crate) use tool_result_replay::limit_tool_result_text;
|
||||
pub use usage::*;
|
||||
|
||||
@@ -0,0 +1,297 @@
|
||||
use super::ProviderType;
|
||||
|
||||
mod usage {
|
||||
use std::ops::AddAssign;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::ProviderType;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub input_tokens: Option<u64>,
|
||||
pub output_tokens: Option<u64>,
|
||||
pub total_tokens: Option<u64>,
|
||||
pub cache_read_tokens: Option<u64>,
|
||||
pub cache_write_tokens: Option<u64>,
|
||||
pub reasoning_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
impl Usage {
|
||||
/// Returns the provider-visible input context without counting cached tokens twice.
|
||||
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
|
||||
let input = self.input_tokens?;
|
||||
match provider {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => Some(input),
|
||||
ProviderType::Anthropic => input
|
||||
.checked_add(self.cache_read_tokens.unwrap_or_default())?
|
||||
.checked_add(self.cache_write_tokens.unwrap_or_default()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AddAssign for Usage {
|
||||
fn add_assign(&mut self, rhs: Self) {
|
||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
||||
self.output_tokens = sum(self.output_tokens, rhs.output_tokens);
|
||||
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
||||
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
|
||||
self.cache_write_tokens = sum(self.cache_write_tokens, rhs.cache_write_tokens);
|
||||
self.reasoning_tokens = sum(self.reasoning_tokens, rhs.reasoning_tokens);
|
||||
}
|
||||
}
|
||||
|
||||
fn sum(left: Option<u64>, right: Option<u64>) -> Option<u64> {
|
||||
left?.checked_add(right?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Usage;
|
||||
use crate::model::ProviderType;
|
||||
|
||||
#[test]
|
||||
fn openai_context_input_does_not_double_count_cached_tokens() {
|
||||
let usage = Usage {
|
||||
input_tokens: Some(140_649),
|
||||
cache_read_tokens: Some(120_000),
|
||||
cache_write_tokens: Some(10_000),
|
||||
..Usage::default()
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::OpenAiResponses),
|
||||
Some(140_649)
|
||||
);
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::OpenAiChat),
|
||||
Some(140_649)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_context_input_includes_disjoint_cache_tokens() {
|
||||
let usage = Usage {
|
||||
input_tokens: Some(10_649),
|
||||
cache_read_tokens: Some(120_000),
|
||||
cache_write_tokens: Some(10_000),
|
||||
..Usage::default()
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::Anthropic),
|
||||
Some(140_649)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn turn_total_only_reports_fields_known_for_every_cycle() {
|
||||
let mut total = Usage {
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(2),
|
||||
total_tokens: Some(12),
|
||||
cache_read_tokens: None,
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: Some(1),
|
||||
};
|
||||
total += Usage {
|
||||
input_tokens: Some(20),
|
||||
output_tokens: Some(3),
|
||||
total_tokens: Some(23),
|
||||
cache_read_tokens: Some(8),
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: None,
|
||||
};
|
||||
assert_eq!(total.input_tokens, Some(30));
|
||||
assert_eq!(total.output_tokens, Some(5));
|
||||
assert_eq!(total.total_tokens, Some(35));
|
||||
assert_eq!(total.cache_read_tokens, None);
|
||||
assert_eq!(total.cache_write_tokens, None);
|
||||
assert_eq!(total.reasoning_tokens, None);
|
||||
}
|
||||
}
|
||||
}
|
||||
pub use usage::*;
|
||||
|
||||
mod llm_call {
|
||||
use serde::Serialize;
|
||||
|
||||
use super::{ProviderType, Usage};
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct NewLlmCall {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub provider_url: String,
|
||||
pub request_type: ProviderType,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: bool,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, PartialEq, Eq)]
|
||||
pub(crate) struct LlmCallUsageAnchor {
|
||||
pub request_type: ProviderType,
|
||||
pub usage: Usage,
|
||||
pub message_count: usize,
|
||||
pub tool_count: usize,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallSummary {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: i64,
|
||||
pub model_hash: Option<String>,
|
||||
pub provider_type: String,
|
||||
pub provider_url: String,
|
||||
pub request_type: String,
|
||||
pub request_url: String,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub fast: Option<bool>,
|
||||
pub status: String,
|
||||
pub finish_reason: Option<String>,
|
||||
pub created_at_ms: i64,
|
||||
pub request_started_at_ms: Option<i64>,
|
||||
pub response_headers_at_ms: Option<i64>,
|
||||
pub first_event_at_ms: Option<i64>,
|
||||
pub first_text_at_ms: Option<i64>,
|
||||
pub first_valid_response_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub queue_ms: Option<i64>,
|
||||
pub ttfb_ms: Option<i64>,
|
||||
pub ttft_ms: Option<i64>,
|
||||
pub ttfr_ms: Option<i64>,
|
||||
pub duration_ms: Option<i64>,
|
||||
pub input_tokens: Option<i64>,
|
||||
pub output_tokens: Option<i64>,
|
||||
pub total_tokens: Option<i64>,
|
||||
pub cache_read_tokens: Option<i64>,
|
||||
pub cache_write_tokens: Option<i64>,
|
||||
pub reasoning_tokens: Option<i64>,
|
||||
pub usage: Option<serde_json::Value>,
|
||||
pub message_count: i64,
|
||||
pub tool_count: i64,
|
||||
pub request_bytes: Option<i64>,
|
||||
pub response_bytes: i64,
|
||||
pub stream_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub error_kind: Option<String>,
|
||||
pub error_message: Option<String>,
|
||||
pub detailed: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallRequest {
|
||||
pub headers: serde_json::Value,
|
||||
pub body: serde_json::Value,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct LlmCallResponseChunk {
|
||||
pub seq: i64,
|
||||
pub received_offset_ms: i64,
|
||||
pub data: String,
|
||||
pub byte_count: i64,
|
||||
}
|
||||
}
|
||||
pub use llm_call::*;
|
||||
|
||||
mod cursor_trace {
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct CursorRunTraceSummary {
|
||||
pub request_id: String,
|
||||
pub conversation_id: Option<String>,
|
||||
pub route: String,
|
||||
pub model_id: Option<String>,
|
||||
pub status: String,
|
||||
pub request_bytes: i64,
|
||||
pub response_bytes: i64,
|
||||
pub response_event_count: i64,
|
||||
pub http_status: Option<i64>,
|
||||
pub received_at_ms: i64,
|
||||
pub first_response_at_ms: Option<i64>,
|
||||
pub finished_at_ms: Option<i64>,
|
||||
pub error_message: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct CursorRunTraceArtifact {
|
||||
pub seq: i64,
|
||||
pub artifact_type: String,
|
||||
pub source: String,
|
||||
pub metadata: serde_json::Value,
|
||||
pub created_at_ms: i64,
|
||||
pub data: Vec<u8>,
|
||||
}
|
||||
}
|
||||
pub use cursor_trace::*;
|
||||
|
||||
mod overview {
|
||||
//! Read-only usage aggregates rendered by the desktop overview page.
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct OverviewMetrics {
|
||||
pub llm_calls: i64,
|
||||
pub successful_calls: i64,
|
||||
pub failed_calls: i64,
|
||||
pub token_usage: i64,
|
||||
pub prompt_tokens: i64,
|
||||
pub input_tokens: i64,
|
||||
pub cache_read_tokens: i64,
|
||||
pub cache_write_tokens: i64,
|
||||
pub output_tokens: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum TokenUsageGranularity {
|
||||
Minute,
|
||||
Hour,
|
||||
#[default]
|
||||
Day,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct TokenUsageBucket {
|
||||
pub bucket_start_ms: i64,
|
||||
pub input_tokens: i64,
|
||||
pub cache_read_tokens: i64,
|
||||
pub cache_write_tokens: i64,
|
||||
pub output_tokens: i64,
|
||||
}
|
||||
|
||||
impl TokenUsageBucket {
|
||||
pub fn total_tokens(&self) -> i64 {
|
||||
self.input_tokens
|
||||
.saturating_add(self.cache_read_tokens)
|
||||
.saturating_add(self.cache_write_tokens)
|
||||
.saturating_add(self.output_tokens)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct Overview {
|
||||
pub metrics: OverviewMetrics,
|
||||
pub token_usage_granularity: TokenUsageGranularity,
|
||||
pub token_usage_series: Vec<TokenUsageBucket>,
|
||||
}
|
||||
}
|
||||
pub use overview::*;
|
||||
@@ -1,50 +0,0 @@
|
||||
//! Read-only usage aggregates rendered by the desktop overview page.
|
||||
|
||||
use serde::Serialize;
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct OverviewMetrics {
|
||||
pub llm_calls: i64,
|
||||
pub successful_calls: i64,
|
||||
pub failed_calls: i64,
|
||||
pub token_usage: i64,
|
||||
pub prompt_tokens: i64,
|
||||
pub input_tokens: i64,
|
||||
pub cache_read_tokens: i64,
|
||||
pub cache_write_tokens: i64,
|
||||
pub output_tokens: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum TokenUsageGranularity {
|
||||
Minute,
|
||||
Hour,
|
||||
#[default]
|
||||
Day,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct TokenUsageBucket {
|
||||
pub bucket_start_ms: i64,
|
||||
pub input_tokens: i64,
|
||||
pub cache_read_tokens: i64,
|
||||
pub cache_write_tokens: i64,
|
||||
pub output_tokens: i64,
|
||||
}
|
||||
|
||||
impl TokenUsageBucket {
|
||||
pub fn total_tokens(&self) -> i64 {
|
||||
self.input_tokens
|
||||
.saturating_add(self.cache_read_tokens)
|
||||
.saturating_add(self.cache_write_tokens)
|
||||
.saturating_add(self.output_tokens)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, PartialEq, Eq, Serialize)]
|
||||
pub struct Overview {
|
||||
pub metrics: OverviewMetrics,
|
||||
pub token_usage_granularity: TokenUsageGranularity,
|
||||
pub token_usage_series: Vec<TokenUsageBucket>,
|
||||
}
|
||||
@@ -1,109 +0,0 @@
|
||||
use std::ops::AddAssign;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::ProviderType;
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct Usage {
|
||||
pub input_tokens: Option<u64>,
|
||||
pub output_tokens: Option<u64>,
|
||||
pub total_tokens: Option<u64>,
|
||||
pub cache_read_tokens: Option<u64>,
|
||||
pub cache_write_tokens: Option<u64>,
|
||||
pub reasoning_tokens: Option<u64>,
|
||||
}
|
||||
|
||||
impl Usage {
|
||||
/// Returns the provider-visible input context without counting cached tokens twice.
|
||||
pub(crate) fn context_input_tokens(self, provider: ProviderType) -> Option<u64> {
|
||||
let input = self.input_tokens?;
|
||||
match provider {
|
||||
ProviderType::OpenAiChat | ProviderType::OpenAiResponses => Some(input),
|
||||
ProviderType::Anthropic => input
|
||||
.checked_add(self.cache_read_tokens.unwrap_or_default())?
|
||||
.checked_add(self.cache_write_tokens.unwrap_or_default()),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl AddAssign for Usage {
|
||||
fn add_assign(&mut self, rhs: Self) {
|
||||
self.input_tokens = sum(self.input_tokens, rhs.input_tokens);
|
||||
self.output_tokens = sum(self.output_tokens, rhs.output_tokens);
|
||||
self.total_tokens = sum(self.total_tokens, rhs.total_tokens);
|
||||
self.cache_read_tokens = sum(self.cache_read_tokens, rhs.cache_read_tokens);
|
||||
self.cache_write_tokens = sum(self.cache_write_tokens, rhs.cache_write_tokens);
|
||||
self.reasoning_tokens = sum(self.reasoning_tokens, rhs.reasoning_tokens);
|
||||
}
|
||||
}
|
||||
|
||||
fn sum(left: Option<u64>, right: Option<u64>) -> Option<u64> {
|
||||
left?.checked_add(right?)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::Usage;
|
||||
use crate::model::ProviderType;
|
||||
|
||||
#[test]
|
||||
fn openai_context_input_does_not_double_count_cached_tokens() {
|
||||
let usage = Usage {
|
||||
input_tokens: Some(140_649),
|
||||
cache_read_tokens: Some(120_000),
|
||||
cache_write_tokens: Some(10_000),
|
||||
..Usage::default()
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::OpenAiResponses),
|
||||
Some(140_649)
|
||||
);
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::OpenAiChat),
|
||||
Some(140_649)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn anthropic_context_input_includes_disjoint_cache_tokens() {
|
||||
let usage = Usage {
|
||||
input_tokens: Some(10_649),
|
||||
cache_read_tokens: Some(120_000),
|
||||
cache_write_tokens: Some(10_000),
|
||||
..Usage::default()
|
||||
};
|
||||
|
||||
assert_eq!(
|
||||
usage.context_input_tokens(ProviderType::Anthropic),
|
||||
Some(140_649)
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn turn_total_only_reports_fields_known_for_every_cycle() {
|
||||
let mut total = Usage {
|
||||
input_tokens: Some(10),
|
||||
output_tokens: Some(2),
|
||||
total_tokens: Some(12),
|
||||
cache_read_tokens: None,
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: Some(1),
|
||||
};
|
||||
total += Usage {
|
||||
input_tokens: Some(20),
|
||||
output_tokens: Some(3),
|
||||
total_tokens: Some(23),
|
||||
cache_read_tokens: Some(8),
|
||||
cache_write_tokens: None,
|
||||
reasoning_tokens: None,
|
||||
};
|
||||
assert_eq!(total.input_tokens, Some(30));
|
||||
assert_eq!(total.output_tokens, Some(5));
|
||||
assert_eq!(total.total_tokens, Some(35));
|
||||
assert_eq!(total.cache_read_tokens, None);
|
||||
assert_eq!(total.cache_write_tokens, None);
|
||||
assert_eq!(total.reasoning_tokens, None);
|
||||
}
|
||||
}
|
||||
@@ -1,56 +0,0 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{ClientCommand, ClientPort},
|
||||
model::PreparedRun,
|
||||
provider::Provider,
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::{RunEngine, RunOutcome, RunRegistry};
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RunActor {
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
registry: RunRegistry,
|
||||
}
|
||||
|
||||
impl RunActor {
|
||||
pub fn new(store: Store, provider: Arc<dyn Provider>, registry: RunRegistry) -> Self {
|
||||
Self {
|
||||
store,
|
||||
provider,
|
||||
registry,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn spawn(
|
||||
&self,
|
||||
prepared: PreparedRun,
|
||||
client: ClientPort,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
cancellation: CancellationToken,
|
||||
) -> tokio::task::JoinHandle<RunOutcome> {
|
||||
let run_id = prepared.run_id.clone();
|
||||
let conversation_id = prepared.conversation_id.clone();
|
||||
self.registry
|
||||
.activate(
|
||||
conversation_id.clone(),
|
||||
run_id.clone(),
|
||||
cancellation.clone(),
|
||||
commands,
|
||||
)
|
||||
.await;
|
||||
let actor = self.clone();
|
||||
tokio::spawn(async move {
|
||||
let outcome = RunEngine::new(actor.store, actor.provider)
|
||||
.run(prepared, client, cancellation)
|
||||
.await;
|
||||
actor.registry.release(&conversation_id, &run_id).await;
|
||||
outcome
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -4,10 +4,6 @@ use std::sync::Arc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{
|
||||
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
|
||||
StateCommitted,
|
||||
},
|
||||
model::{
|
||||
CanonicalMessage, MessageContent, Origin, PreparedRun, Role, RunAction, ToolRoundAssistant,
|
||||
ToolRoundId, Usage,
|
||||
@@ -16,7 +12,10 @@ use crate::{
|
||||
store::{RunStatus, Store},
|
||||
};
|
||||
|
||||
use super::{consume_model_cycle, ModelCycleFailure, RunFailure, RunOutcome};
|
||||
use super::{
|
||||
consume_model_cycle, ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause,
|
||||
MessageInsertion, ModelCycleFailure, RunFailure, RunOutcome, StateCommitted,
|
||||
};
|
||||
|
||||
const COMPACTION_RESERVE_TOKENS: u64 = 10_000;
|
||||
const COMPACTION_OUTPUT_TOKENS: u64 = 4_096;
|
||||
|
||||
@@ -1,45 +0,0 @@
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RunFailure {
|
||||
Protocol(String),
|
||||
Provider(String),
|
||||
Store(String),
|
||||
Client(String),
|
||||
}
|
||||
|
||||
impl RunFailure {
|
||||
pub fn category(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Protocol(_) => "protocol",
|
||||
Self::Provider(_) => "provider",
|
||||
Self::Store(_) => "store",
|
||||
Self::Client(_) => "client",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<crate::Error> for RunFailure {
|
||||
fn from(error: crate::Error) -> Self {
|
||||
use crate::Error;
|
||||
match error {
|
||||
Error::Protocol(message) | Error::Config(message) => Self::Protocol(message),
|
||||
Error::Provider(message) => Self::Provider(message),
|
||||
Error::Store(message) => Self::Store(message),
|
||||
Error::Cancelled => Self::Client("run was cancelled".into()),
|
||||
Error::Http(error) => Self::Provider(error.to_string()),
|
||||
Error::Database(error) => Self::Store(error.to_string()),
|
||||
Error::Migration(error) => Self::Store(error.to_string()),
|
||||
Error::Io(error) => Self::Store(error.to_string()),
|
||||
Error::Decode(error) => Self::Protocol(error.to_string()),
|
||||
Error::Encode(error) => Self::Protocol(error.to_string()),
|
||||
Error::Json(error) => Self::Protocol(error.to_string()),
|
||||
Error::RunNotFound(run_id) => Self::Store(format!("run not found: {run_id}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RunOutcome {
|
||||
Completed,
|
||||
Cancelled,
|
||||
Failed(RunFailure),
|
||||
}
|
||||
@@ -1,12 +1,10 @@
|
||||
mod actor;
|
||||
mod engine;
|
||||
mod lifecycle;
|
||||
mod model_cycle;
|
||||
mod registry;
|
||||
mod port;
|
||||
mod runtime;
|
||||
mod tool_round;
|
||||
|
||||
pub use actor::*;
|
||||
pub use engine::*;
|
||||
pub use lifecycle::*;
|
||||
pub use model_cycle::*;
|
||||
pub use registry::*;
|
||||
pub use port::*;
|
||||
pub use runtime::*;
|
||||
|
||||
@@ -5,12 +5,11 @@ use tokio::sync::mpsc;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::ClientEvent,
|
||||
model::{ProviderReplayState, ToolCall, Usage},
|
||||
provider::{FinishReason, ModelEvent, ProviderStream},
|
||||
};
|
||||
|
||||
use super::RunFailure;
|
||||
use super::{ClientEvent, RunFailure};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq)]
|
||||
pub struct ModelCycleResult {
|
||||
|
||||
@@ -1,9 +1,28 @@
|
||||
use std::time::Duration;
|
||||
|
||||
use tokio::sync::oneshot;
|
||||
use tokio::sync::{mpsc, oneshot};
|
||||
|
||||
use crate::model::{RevisionId, ToolCall, ToolRoundId, Usage};
|
||||
use crate::run::RunOutcome;
|
||||
use crate::model::{
|
||||
CanonicalMessage, RevisionId, RuntimeEvent, ToolCall, ToolResult, ToolRoundId, Usage,
|
||||
};
|
||||
|
||||
use super::RunOutcome;
|
||||
|
||||
#[derive(Debug)]
|
||||
pub struct MessageInsertion {
|
||||
pub messages: Vec<CanonicalMessage>,
|
||||
pub delivered: oneshot::Sender<()>,
|
||||
}
|
||||
|
||||
#[derive(Debug)]
|
||||
pub enum ClientCommand {
|
||||
ToolResult(ToolResult),
|
||||
InterruptWithMessage(CanonicalMessage),
|
||||
RuntimeEvent(RuntimeEvent),
|
||||
InsertMessages(MessageInsertion),
|
||||
ClientClosed { error: String },
|
||||
Cancel,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum CommitCause {
|
||||
@@ -79,3 +98,28 @@ pub enum ClientEvent {
|
||||
StateCommitted(StateCommitted),
|
||||
Ended(RunOutcome),
|
||||
}
|
||||
|
||||
pub struct ClientPort {
|
||||
pub commands: mpsc::Receiver<ClientCommand>,
|
||||
pub events: mpsc::Sender<ClientEvent>,
|
||||
}
|
||||
|
||||
pub struct ClientSession {
|
||||
pub commands: mpsc::Sender<ClientCommand>,
|
||||
pub events: mpsc::Receiver<ClientEvent>,
|
||||
}
|
||||
|
||||
pub fn session(capacity: usize) -> (ClientPort, ClientSession) {
|
||||
let (commands_tx, commands_rx) = mpsc::channel(capacity);
|
||||
let (events_tx, events_rx) = mpsc::channel(capacity);
|
||||
(
|
||||
ClientPort {
|
||||
commands: commands_rx,
|
||||
events: events_tx,
|
||||
},
|
||||
ClientSession {
|
||||
commands: commands_tx,
|
||||
events: events_rx,
|
||||
},
|
||||
)
|
||||
}
|
||||
@@ -1,90 +0,0 @@
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{ClientCommand, MessageInsertion},
|
||||
model::{CanonicalMessage, ConversationId, RunId},
|
||||
};
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct RunRegistry {
|
||||
active: Arc<Mutex<HashMap<ConversationId, ActiveRun>>>,
|
||||
}
|
||||
|
||||
struct ActiveRun {
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
}
|
||||
|
||||
impl RunRegistry {
|
||||
pub async fn activate(
|
||||
&self,
|
||||
conversation_id: ConversationId,
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
) {
|
||||
let previous = self.active.lock().await.insert(
|
||||
conversation_id,
|
||||
ActiveRun {
|
||||
run_id: run_id.clone(),
|
||||
cancellation,
|
||||
commands,
|
||||
},
|
||||
);
|
||||
if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) {
|
||||
previous.cancellation.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn insert_messages(
|
||||
&self,
|
||||
conversation_id: &ConversationId,
|
||||
messages: Vec<CanonicalMessage>,
|
||||
) -> bool {
|
||||
if messages.is_empty() {
|
||||
return true;
|
||||
}
|
||||
let commands = self
|
||||
.active
|
||||
.lock()
|
||||
.await
|
||||
.get(conversation_id)
|
||||
.map(|run| run.commands.clone());
|
||||
let Some(commands) = commands else {
|
||||
return false;
|
||||
};
|
||||
let (delivered, delivery) = tokio::sync::oneshot::channel();
|
||||
if commands
|
||||
.send(ClientCommand::InsertMessages(MessageInsertion {
|
||||
messages,
|
||||
delivered,
|
||||
}))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
delivery.await.is_ok()
|
||||
}
|
||||
|
||||
pub async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) {
|
||||
let mut active = self.active.lock().await;
|
||||
if active
|
||||
.get(conversation_id)
|
||||
.is_some_and(|current| ¤t.run_id == run_id)
|
||||
{
|
||||
active.remove(conversation_id);
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn shutdown(&self) {
|
||||
let active = std::mem::take(&mut *self.active.lock().await);
|
||||
for run in active.into_values() {
|
||||
run.cancellation.cancel();
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,183 @@
|
||||
use std::{collections::HashMap, sync::Arc};
|
||||
|
||||
use tokio::sync::Mutex;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
model::{CanonicalMessage, ConversationId, PreparedRun, RunId},
|
||||
provider::Provider,
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::{ClientCommand, ClientPort, MessageInsertion, RunEngine};
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RunFailure {
|
||||
Protocol(String),
|
||||
Provider(String),
|
||||
Store(String),
|
||||
Client(String),
|
||||
}
|
||||
|
||||
impl RunFailure {
|
||||
pub fn category(&self) -> &'static str {
|
||||
match self {
|
||||
Self::Protocol(_) => "protocol",
|
||||
Self::Provider(_) => "provider",
|
||||
Self::Store(_) => "store",
|
||||
Self::Client(_) => "client",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl From<crate::Error> for RunFailure {
|
||||
fn from(error: crate::Error) -> Self {
|
||||
use crate::Error;
|
||||
match error {
|
||||
Error::Protocol(message) | Error::Config(message) => Self::Protocol(message),
|
||||
Error::Provider(message) => Self::Provider(message),
|
||||
Error::Store(message) => Self::Store(message),
|
||||
Error::Cancelled => Self::Client("run was cancelled".into()),
|
||||
Error::Http(error) => Self::Provider(error.to_string()),
|
||||
Error::Database(error) => Self::Store(error.to_string()),
|
||||
Error::Migration(error) => Self::Store(error.to_string()),
|
||||
Error::Io(error) => Self::Store(error.to_string()),
|
||||
Error::Decode(error) => Self::Protocol(error.to_string()),
|
||||
Error::Encode(error) => Self::Protocol(error.to_string()),
|
||||
Error::Json(error) => Self::Protocol(error.to_string()),
|
||||
Error::RunNotFound(run_id) => Self::Store(format!("run not found: {run_id}")),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, PartialEq, Eq)]
|
||||
pub enum RunOutcome {
|
||||
Completed,
|
||||
Cancelled,
|
||||
Failed(RunFailure),
|
||||
}
|
||||
|
||||
#[derive(Clone, Default)]
|
||||
pub struct RunRegistry {
|
||||
active: Arc<Mutex<HashMap<ConversationId, ActiveRun>>>,
|
||||
}
|
||||
|
||||
struct ActiveRun {
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
}
|
||||
|
||||
impl RunRegistry {
|
||||
pub async fn activate(
|
||||
&self,
|
||||
conversation_id: ConversationId,
|
||||
run_id: RunId,
|
||||
cancellation: CancellationToken,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
) {
|
||||
let previous = self.active.lock().await.insert(
|
||||
conversation_id,
|
||||
ActiveRun {
|
||||
run_id: run_id.clone(),
|
||||
cancellation,
|
||||
commands,
|
||||
},
|
||||
);
|
||||
if let Some(previous) = previous.filter(|previous| previous.run_id != run_id) {
|
||||
previous.cancellation.cancel();
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn insert_messages(
|
||||
&self,
|
||||
conversation_id: &ConversationId,
|
||||
messages: Vec<CanonicalMessage>,
|
||||
) -> bool {
|
||||
if messages.is_empty() {
|
||||
return true;
|
||||
}
|
||||
let commands = self
|
||||
.active
|
||||
.lock()
|
||||
.await
|
||||
.get(conversation_id)
|
||||
.map(|run| run.commands.clone());
|
||||
let Some(commands) = commands else {
|
||||
return false;
|
||||
};
|
||||
let (delivered, delivery) = tokio::sync::oneshot::channel();
|
||||
if commands
|
||||
.send(ClientCommand::InsertMessages(MessageInsertion {
|
||||
messages,
|
||||
delivered,
|
||||
}))
|
||||
.await
|
||||
.is_err()
|
||||
{
|
||||
return false;
|
||||
}
|
||||
delivery.await.is_ok()
|
||||
}
|
||||
|
||||
pub async fn release(&self, conversation_id: &ConversationId, run_id: &RunId) {
|
||||
let mut active = self.active.lock().await;
|
||||
if active
|
||||
.get(conversation_id)
|
||||
.is_some_and(|current| ¤t.run_id == run_id)
|
||||
{
|
||||
active.remove(conversation_id);
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn shutdown(&self) {
|
||||
let active = std::mem::take(&mut *self.active.lock().await);
|
||||
for run in active.into_values() {
|
||||
run.cancellation.cancel();
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone)]
|
||||
pub struct RunActor {
|
||||
store: Store,
|
||||
provider: Arc<dyn Provider>,
|
||||
registry: RunRegistry,
|
||||
}
|
||||
|
||||
impl RunActor {
|
||||
pub fn new(store: Store, provider: Arc<dyn Provider>, registry: RunRegistry) -> Self {
|
||||
Self {
|
||||
store,
|
||||
provider,
|
||||
registry,
|
||||
}
|
||||
}
|
||||
|
||||
pub async fn spawn(
|
||||
&self,
|
||||
prepared: PreparedRun,
|
||||
client: ClientPort,
|
||||
commands: tokio::sync::mpsc::Sender<ClientCommand>,
|
||||
cancellation: CancellationToken,
|
||||
) -> tokio::task::JoinHandle<RunOutcome> {
|
||||
let run_id = prepared.run_id.clone();
|
||||
let conversation_id = prepared.conversation_id.clone();
|
||||
self.registry
|
||||
.activate(
|
||||
conversation_id.clone(),
|
||||
run_id.clone(),
|
||||
cancellation.clone(),
|
||||
commands,
|
||||
)
|
||||
.await;
|
||||
let actor = self.clone();
|
||||
tokio::spawn(async move {
|
||||
let outcome = RunEngine::new(actor.store, actor.provider)
|
||||
.run(prepared, client, cancellation)
|
||||
.await;
|
||||
actor.registry.release(&conversation_id, &run_id).await;
|
||||
outcome
|
||||
})
|
||||
}
|
||||
}
|
||||
@@ -3,15 +3,14 @@ use std::collections::HashSet;
|
||||
use tokio_util::sync::CancellationToken;
|
||||
|
||||
use crate::{
|
||||
client::{
|
||||
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
|
||||
StateCommitted,
|
||||
},
|
||||
model::{PreparedRun, RevisionId, ToolCall, ToolResult, ToolRoundAssistant, ToolRoundId},
|
||||
store::Store,
|
||||
};
|
||||
|
||||
use super::{RunFailure, RunOutcome};
|
||||
use super::{
|
||||
ClientCommand, ClientEvent, ClientPort, CommitBarrier, CommitCause, MessageInsertion,
|
||||
RunFailure, RunOutcome, StateCommitted,
|
||||
};
|
||||
|
||||
pub(super) struct ToolRound {
|
||||
pub id: ToolRoundId,
|
||||
|
||||
@@ -2,7 +2,9 @@ mod catalog;
|
||||
mod engine;
|
||||
mod federation;
|
||||
mod fetch;
|
||||
mod semble;
|
||||
|
||||
pub use engine::{HtmlEngine, JsonEngine, SearchEngine, SearchHit};
|
||||
pub use federation::{SearchError, WebSearch};
|
||||
pub use fetch::{FetchError, FetchedPage, WebFetch};
|
||||
pub(crate) use semble::execute as execute_semble;
|
||||
@@ -0,0 +1,172 @@
|
||||
use std::sync::Arc;
|
||||
|
||||
use semble_core::{ContentType, FindRelatedRequest, SearchEngine, SearchRequest, SembleConfig};
|
||||
use serde::Deserialize;
|
||||
use serde_json::Value;
|
||||
use tokio::sync::OnceCell;
|
||||
|
||||
use crate::{store::Store, Error, Result};
|
||||
|
||||
static ENGINE: OnceCell<Arc<SearchEngine>> = OnceCell::const_new();
|
||||
|
||||
#[derive(Clone, Copy, Debug, Default, Deserialize)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
enum ContentSelection {
|
||||
#[default]
|
||||
Code,
|
||||
Docs,
|
||||
Config,
|
||||
All,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct SearchArguments {
|
||||
query: String,
|
||||
repo: String,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
#[derive(Debug, Deserialize)]
|
||||
struct FindRelatedArguments {
|
||||
repo: String,
|
||||
file_path: String,
|
||||
line: usize,
|
||||
#[serde(default = "default_top_k")]
|
||||
top_k: usize,
|
||||
#[serde(default = "default_snippet_lines")]
|
||||
max_snippet_lines: Option<usize>,
|
||||
#[serde(default)]
|
||||
content: ContentSelection,
|
||||
}
|
||||
|
||||
enum Operation {
|
||||
Search(SearchArguments),
|
||||
FindRelated(FindRelatedArguments),
|
||||
}
|
||||
|
||||
pub(crate) async fn execute(
|
||||
tool_name: &str,
|
||||
arguments: Value,
|
||||
store: Option<Store>,
|
||||
) -> std::result::Result<Value, String> {
|
||||
let operation = match tool_name {
|
||||
"semblesearch" => {
|
||||
Operation::Search(serde_json::from_value(arguments).map_err(|error| error.to_string())?)
|
||||
}
|
||||
"semblefindrelated" => Operation::FindRelated(
|
||||
serde_json::from_value(arguments).map_err(|error| error.to_string())?,
|
||||
),
|
||||
_ => return Err(format!("unsupported Semble tool: {tool_name}")),
|
||||
};
|
||||
let engine = engine(store).await.map_err(|error| error.to_string())?;
|
||||
tokio::task::spawn_blocking(move || match operation {
|
||||
Operation::Search(arguments) => engine
|
||||
.search(SearchRequest {
|
||||
query: arguments.query,
|
||||
repo: arguments.repo.into(),
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
Operation::FindRelated(arguments) => engine
|
||||
.find_related(FindRelatedRequest {
|
||||
repo: arguments.repo.into(),
|
||||
file_path: arguments.file_path,
|
||||
line: arguments.line,
|
||||
top_k: arguments.top_k,
|
||||
max_snippet_lines: arguments.max_snippet_lines,
|
||||
content: content(arguments.content),
|
||||
})
|
||||
.and_then(json_value),
|
||||
})
|
||||
.await
|
||||
.map_err(|error| format!("Semble search worker failed: {error}"))?
|
||||
.map_err(|error| error.to_string())
|
||||
}
|
||||
|
||||
async fn engine(store: Option<Store>) -> Result<Arc<SearchEngine>> {
|
||||
ENGINE
|
||||
.get_or_try_init(|| async move {
|
||||
let builder = match store {
|
||||
Some(store) => crate::network::blocking_client_builder(&store).await?,
|
||||
None => reqwest::blocking::Client::builder().use_native_tls(),
|
||||
};
|
||||
tokio::task::spawn_blocking(move || {
|
||||
let client = builder.build()?;
|
||||
SearchEngine::load_default_with_client(SembleConfig::default(), &client)
|
||||
.map(Arc::new)
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))
|
||||
})
|
||||
.await
|
||||
.map_err(|error| Error::Config(format!("load Semble search engine: {error}")))?
|
||||
})
|
||||
.await
|
||||
.cloned()
|
||||
}
|
||||
|
||||
fn json_value(response: semble_core::SearchResponse) -> semble_core::Result<Value> {
|
||||
serde_json::to_value(response)
|
||||
.map_err(|error| semble_core::Error::Serialization(error.to_string()))
|
||||
}
|
||||
|
||||
fn content(selection: ContentSelection) -> Vec<ContentType> {
|
||||
match selection {
|
||||
ContentSelection::Code => vec![ContentType::Code],
|
||||
ContentSelection::Docs => vec![ContentType::Docs],
|
||||
ContentSelection::Config => vec![ContentType::Config],
|
||||
ContentSelection::All => vec![ContentType::Code, ContentType::Docs, ContentType::Config],
|
||||
}
|
||||
}
|
||||
|
||||
fn default_top_k() -> usize {
|
||||
5
|
||||
}
|
||||
|
||||
fn default_snippet_lines() -> Option<usize> {
|
||||
Some(10)
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use serde_json::json;
|
||||
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn search_arguments_use_code_search_defaults() {
|
||||
let arguments: SearchArguments = serde_json::from_value(json!({
|
||||
"query": "request persistence",
|
||||
"repo": "/tmp/repo"
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.top_k, 5);
|
||||
assert_eq!(arguments.max_snippet_lines, Some(10));
|
||||
assert!(matches!(arguments.content, ContentSelection::Code));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn find_related_does_not_require_a_ui_description() {
|
||||
let arguments: FindRelatedArguments = serde_json::from_value(json!({
|
||||
"repo": "/tmp/repo",
|
||||
"file_path": "src/auth.ts",
|
||||
"line": 42
|
||||
}))
|
||||
.unwrap();
|
||||
assert_eq!(arguments.file_path, "src/auth.ts");
|
||||
assert_eq!(arguments.line, 42);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn all_content_expands_to_every_indexed_scope() {
|
||||
assert_eq!(
|
||||
content(ContentSelection::All),
|
||||
vec![ContentType::Code, ContentType::Docs, ContentType::Config]
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -6,13 +6,12 @@ mod fixtures;
|
||||
use std::sync::Arc;
|
||||
|
||||
use cursor_server::{
|
||||
client::{session, ClientCommand, ClientEvent, CommitCause},
|
||||
model::{
|
||||
ConversationId, ModelSpec, PreparedRun, PromptSpec, RunAction, RunId, RunKind,
|
||||
ToolDefinition, ToolResult,
|
||||
},
|
||||
provider::{FinishReason, ModelEvent},
|
||||
run::{RunEngine, RunOutcome},
|
||||
run::{session, ClientCommand, ClientEvent, CommitCause, RunEngine, RunOutcome},
|
||||
};
|
||||
use tokio::{sync::oneshot, time::Duration};
|
||||
use tokio_util::sync::CancellationToken;
|
||||
@@ -94,7 +93,7 @@ async fn inserted_messages_wait_for_the_next_model_call_without_interrupting_the
|
||||
let (delivered, mut delivery) = oneshot::channel();
|
||||
commands
|
||||
.send(ClientCommand::InsertMessages(
|
||||
cursor_server::client::MessageInsertion {
|
||||
cursor_server::run::MessageInsertion {
|
||||
messages: vec![cursor_server::model::RuntimeEvent {
|
||||
event_id: "background:finished".into(),
|
||||
text: "background work finished".into(),
|
||||
|
||||
@@ -32,7 +32,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("first"),
|
||||
first.clone(),
|
||||
cursor_server::client::session(1).1.commands,
|
||||
cursor_server::run::session(1).1.commands,
|
||||
)
|
||||
.await;
|
||||
registry
|
||||
@@ -40,7 +40,7 @@ async fn generic_run_registry_cancels_the_previous_client_for_a_conversation() {
|
||||
conversation.clone(),
|
||||
cursor_server::model::RunId::new("second"),
|
||||
second.clone(),
|
||||
cursor_server::client::session(1).1.commands,
|
||||
cursor_server::run::session(1).1.commands,
|
||||
)
|
||||
.await;
|
||||
|
||||
|
||||
@@ -1,5 +1,4 @@
|
||||
use cursor_server::{
|
||||
client::ClientEvent,
|
||||
config::{ProviderConfig, ProviderKind},
|
||||
model::{
|
||||
ContentPart, ModelInvocation, ModelLatency, ModelRequest, ModelSpec, ProjectedContent,
|
||||
@@ -9,7 +8,7 @@ use cursor_server::{
|
||||
FinishReason, ModelEvent, OpenAiChatProvider, OpenAiResponsesProvider, Provider,
|
||||
ProviderStream,
|
||||
},
|
||||
run::{consume_model_cycle, RunFailure},
|
||||
run::{consume_model_cycle, ClientEvent, RunFailure},
|
||||
};
|
||||
use futures_util::{stream, StreamExt};
|
||||
use serde_json::{json, Value};
|
||||
|
||||
@@ -1,5 +1,5 @@
|
||||
use axum::{http::StatusCode, response::IntoResponse, routing::get, Router};
|
||||
use cursor_server::web::{HtmlEngine, JsonEngine, WebSearch};
|
||||
use cursor_server::search::{HtmlEngine, JsonEngine, WebSearch};
|
||||
use tokio::net::TcpListener;
|
||||
|
||||
const RESULT_SELECTOR: &str = ".result";
|
||||
@@ -106,7 +106,7 @@ async fn json_engines_use_declared_result_fields() {
|
||||
#[ignore = "live public search smoke test"]
|
||||
async fn built_in_search_returns_live_results() {
|
||||
let _ = tracing_subscriber::fmt()
|
||||
.with_env_filter("cursor_server::web=debug")
|
||||
.with_env_filter("cursor_server::search=debug")
|
||||
.try_init();
|
||||
let results = WebSearch::built_in()
|
||||
.search("Rust programming language")
|
||||
|
||||
Reference in New Issue
Block a user