mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 20:44:07 +08:00
refactor: rebuild desktop app with Tauri
This commit is contained in:
@@ -0,0 +1,60 @@
|
||||
use std::fmt;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
macro_rules! string_id {
|
||||
($name:ident) => {
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord)]
|
||||
#[serde(transparent)]
|
||||
pub struct $name(pub String);
|
||||
|
||||
impl $name {
|
||||
pub fn new(value: impl Into<String>) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
|
||||
pub fn as_str(&self) -> &str {
|
||||
&self.0
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for $name {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<String> for $name {
|
||||
fn from(value: String) -> Self {
|
||||
Self(value)
|
||||
}
|
||||
}
|
||||
|
||||
impl From<&str> for $name {
|
||||
fn from(value: &str) -> Self {
|
||||
Self(value.into())
|
||||
}
|
||||
}
|
||||
};
|
||||
}
|
||||
|
||||
string_id!(ConversationId);
|
||||
string_id!(RunId);
|
||||
string_id!(ToolRoundId);
|
||||
|
||||
#[derive(Clone, Copy, Debug, Serialize, Deserialize, PartialEq, Eq, Hash, PartialOrd, Ord)]
|
||||
#[serde(transparent)]
|
||||
pub struct RevisionId(pub i64);
|
||||
|
||||
impl fmt::Display for RevisionId {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
self.0.fmt(formatter)
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct Conversation {
|
||||
pub conversation_id: ConversationId,
|
||||
pub current_revision_id: RevisionId,
|
||||
pub active_run_id: Option<RunId>,
|
||||
}
|
||||
@@ -0,0 +1,28 @@
|
||||
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>,
|
||||
}
|
||||
@@ -0,0 +1,25 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{ModelSpec, ProjectedMessage, ToolDefinition};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct PromptSpec {
|
||||
pub instructions: String,
|
||||
pub tools: Vec<ToolDefinition>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ModelRequest {
|
||||
pub prompt: PromptSpec,
|
||||
pub model: ModelSpec,
|
||||
pub history: Vec<ProjectedMessage>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ModelInvocation {
|
||||
pub call_id: String,
|
||||
pub run_id: String,
|
||||
pub conversation_id: String,
|
||||
pub provider_call_index: u64,
|
||||
pub request: ModelRequest,
|
||||
}
|
||||
@@ -0,0 +1,83 @@
|
||||
use serde::Serialize;
|
||||
|
||||
use super::ProviderType;
|
||||
|
||||
#[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, 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 finished_at_ms: Option<i64>,
|
||||
pub queue_ms: Option<i64>,
|
||||
pub ttfb_ms: Option<i64>,
|
||||
pub ttft_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,
|
||||
}
|
||||
@@ -0,0 +1,144 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::{ToolImageReference, ToolRoundId};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum Role {
|
||||
System,
|
||||
User,
|
||||
Assistant,
|
||||
Tool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum Origin {
|
||||
Prompt,
|
||||
User,
|
||||
Runtime,
|
||||
Assistant,
|
||||
Tool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolCallContent {
|
||||
pub index: usize,
|
||||
pub call_id: String,
|
||||
pub name: String,
|
||||
pub arguments: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolResultContent {
|
||||
pub call_id: String,
|
||||
pub name: String,
|
||||
pub content: String,
|
||||
pub is_error: bool,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
pub image: Option<ToolImageReference>,
|
||||
#[serde(skip)]
|
||||
pub provider_parts: Vec<ContentPart>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ProviderReplayState {
|
||||
pub provider_kind: String,
|
||||
pub value: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "type", rename_all = "snake_case")]
|
||||
pub enum ContentPart {
|
||||
Text {
|
||||
text: String,
|
||||
},
|
||||
Image {
|
||||
mime_type: String,
|
||||
#[serde(with = "base64_bytes")]
|
||||
data: Vec<u8>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
#[serde(tag = "kind", rename_all = "snake_case")]
|
||||
pub enum MessageContent {
|
||||
Parts {
|
||||
parts: Vec<ContentPart>,
|
||||
},
|
||||
Assistant {
|
||||
text: String,
|
||||
thinking: String,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
tool_round_id: Option<ToolRoundId>,
|
||||
#[serde(default, skip_serializing_if = "Option::is_none")]
|
||||
replay_state: Option<ProviderReplayState>,
|
||||
tool_calls: Vec<ToolCallContent>,
|
||||
},
|
||||
ToolResult(ToolResultContent),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct CanonicalMessage {
|
||||
pub message_id: String,
|
||||
pub role: Role,
|
||||
pub origin: Origin,
|
||||
pub content: MessageContent,
|
||||
#[serde(skip_serializing_if = "Option::is_none")]
|
||||
pub runtime_event_id: Option<String>,
|
||||
}
|
||||
|
||||
impl CanonicalMessage {
|
||||
pub fn text(
|
||||
message_id: impl Into<String>,
|
||||
role: Role,
|
||||
origin: Origin,
|
||||
text: impl Into<String>,
|
||||
) -> Self {
|
||||
Self {
|
||||
message_id: message_id.into(),
|
||||
role,
|
||||
origin,
|
||||
content: MessageContent::Parts {
|
||||
parts: vec![ContentPart::Text { text: text.into() }],
|
||||
},
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
|
||||
pub fn parts(
|
||||
message_id: impl Into<String>,
|
||||
role: Role,
|
||||
origin: Origin,
|
||||
parts: Vec<ContentPart>,
|
||||
) -> Self {
|
||||
Self {
|
||||
message_id: message_id.into(),
|
||||
role,
|
||||
origin,
|
||||
content: MessageContent::Parts { parts },
|
||||
runtime_event_id: None,
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
mod base64_bytes {
|
||||
use base64::{engine::general_purpose::STANDARD, Engine};
|
||||
use serde::{Deserialize, Deserializer, Serializer};
|
||||
|
||||
pub fn serialize<S>(data: &[u8], serializer: S) -> Result<S::Ok, S::Error>
|
||||
where
|
||||
S: Serializer,
|
||||
{
|
||||
serializer.serialize_str(&STANDARD.encode(data))
|
||||
}
|
||||
|
||||
pub fn deserialize<'de, D>(deserializer: D) -> Result<Vec<u8>, D::Error>
|
||||
where
|
||||
D: Deserializer<'de>,
|
||||
{
|
||||
let encoded = String::deserialize(deserializer)?;
|
||||
STANDARD.decode(encoded).map_err(serde::de::Error::custom)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,27 @@
|
||||
mod conversation;
|
||||
mod cursor_trace;
|
||||
mod inference;
|
||||
mod llm_call;
|
||||
mod message;
|
||||
mod model_spec;
|
||||
mod overview;
|
||||
mod projection;
|
||||
mod provider;
|
||||
mod run;
|
||||
mod runtime_tag;
|
||||
mod tool;
|
||||
mod usage;
|
||||
|
||||
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 projection::*;
|
||||
pub use provider::*;
|
||||
pub use run::*;
|
||||
pub use runtime_tag::*;
|
||||
pub use tool::*;
|
||||
pub use usage::*;
|
||||
@@ -0,0 +1,44 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct ReasoningSpec {
|
||||
pub enabled: bool,
|
||||
pub effort: Option<String>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Default, Serialize, Deserialize, PartialEq, Eq)]
|
||||
#[serde(rename_all = "snake_case")]
|
||||
pub enum ModelLatency {
|
||||
#[default]
|
||||
Standard,
|
||||
Fast,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ModelSpec {
|
||||
pub model_id: String,
|
||||
pub display_name: Option<String>,
|
||||
pub reasoning: ReasoningSpec,
|
||||
pub latency: ModelLatency,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub context_window_tokens: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub supports_image_generation: bool,
|
||||
#[serde(default)]
|
||||
pub extra_params: serde_json::Value,
|
||||
}
|
||||
|
||||
impl ModelSpec {
|
||||
pub fn new(model_id: impl Into<String>) -> Self {
|
||||
Self {
|
||||
model_id: model_id.into(),
|
||||
display_name: None,
|
||||
reasoning: ReasoningSpec::default(),
|
||||
latency: ModelLatency::Standard,
|
||||
max_output_tokens: None,
|
||||
context_window_tokens: None,
|
||||
supports_image_generation: false,
|
||||
extra_params: serde_json::json!({}),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,50 @@
|
||||
//! 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>,
|
||||
}
|
||||
@@ -0,0 +1,169 @@
|
||||
use std::collections::HashSet;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
use super::{
|
||||
CanonicalMessage, ContentPart, MessageContent, ProviderReplayState, Role, ToolCallContent,
|
||||
ToolResultContent,
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub enum ProjectedContent {
|
||||
Parts(Vec<ContentPart>),
|
||||
Assistant {
|
||||
text: String,
|
||||
thinking: String,
|
||||
replay_state: Option<ProviderReplayState>,
|
||||
calls: Vec<ToolCallContent>,
|
||||
},
|
||||
ToolResult(ToolResultContent),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ProjectedMessage {
|
||||
pub message_id: String,
|
||||
pub role: Role,
|
||||
pub content: ProjectedContent,
|
||||
}
|
||||
|
||||
pub fn project_messages(messages: &[CanonicalMessage]) -> Result<Vec<ProjectedMessage>> {
|
||||
let mut projected = Vec::new();
|
||||
let mut index = 0;
|
||||
while index < messages.len() {
|
||||
if let Some((group, next)) = project_tool_round(messages, index)? {
|
||||
projected.extend(group);
|
||||
index = next;
|
||||
} else {
|
||||
projected.push(project_message(&messages[index]));
|
||||
index += 1;
|
||||
}
|
||||
}
|
||||
Ok(projected)
|
||||
}
|
||||
|
||||
fn project_tool_round(
|
||||
messages: &[CanonicalMessage],
|
||||
start: usize,
|
||||
) -> Result<Option<(Vec<ProjectedMessage>, usize)>> {
|
||||
let MessageContent::Assistant {
|
||||
tool_round_id: Some(group_id),
|
||||
tool_calls,
|
||||
..
|
||||
} = &messages[start].content
|
||||
else {
|
||||
return Ok(None);
|
||||
};
|
||||
if tool_calls.is_empty() {
|
||||
return Ok(None);
|
||||
}
|
||||
|
||||
let mut cursor = start;
|
||||
let mut text = String::new();
|
||||
let mut thinking = String::new();
|
||||
let mut replay_state = None;
|
||||
let mut calls = Vec::new();
|
||||
let mut results = Vec::new();
|
||||
let mut result_ids = HashSet::new();
|
||||
|
||||
while cursor < messages.len() {
|
||||
let MessageContent::Assistant {
|
||||
text: part_text,
|
||||
thinking: part_thinking,
|
||||
tool_round_id: Some(candidate_group),
|
||||
replay_state: part_replay,
|
||||
tool_calls: part_calls,
|
||||
} = &messages[cursor].content
|
||||
else {
|
||||
break;
|
||||
};
|
||||
if candidate_group != group_id || part_calls.is_empty() {
|
||||
break;
|
||||
}
|
||||
text.push_str(part_text);
|
||||
thinking.push_str(part_thinking);
|
||||
if replay_state.is_none() {
|
||||
replay_state = part_replay.clone();
|
||||
} else if part_replay.is_some() {
|
||||
return Err(Error::Protocol(
|
||||
"tool round repeats provider replay state".into(),
|
||||
));
|
||||
}
|
||||
calls.extend(part_calls.iter().cloned());
|
||||
cursor += 1;
|
||||
|
||||
while cursor < messages.len() {
|
||||
let MessageContent::ToolResult(result) = &messages[cursor].content else {
|
||||
break;
|
||||
};
|
||||
if !calls.iter().any(|call| call.call_id == result.call_id) {
|
||||
break;
|
||||
}
|
||||
if !result_ids.insert(result.call_id.clone()) {
|
||||
return Err(Error::Protocol(format!(
|
||||
"duplicate tool result call_id: {}",
|
||||
result.call_id
|
||||
)));
|
||||
}
|
||||
results.push((messages[cursor].message_id.clone(), result.clone()));
|
||||
cursor += 1;
|
||||
}
|
||||
}
|
||||
|
||||
calls.sort_by_key(|call| call.index);
|
||||
for call in &calls {
|
||||
if !result_ids.contains(&call.call_id) {
|
||||
return Err(Error::Protocol(format!(
|
||||
"assistant tool call has no result call_id: {}",
|
||||
call.call_id
|
||||
)));
|
||||
}
|
||||
}
|
||||
|
||||
let mut output = Vec::with_capacity(results.len() + 1);
|
||||
output.push(ProjectedMessage {
|
||||
message_id: messages[start].message_id.clone(),
|
||||
role: Role::Assistant,
|
||||
content: ProjectedContent::Assistant {
|
||||
text,
|
||||
thinking,
|
||||
replay_state,
|
||||
calls,
|
||||
},
|
||||
});
|
||||
output.extend(
|
||||
results
|
||||
.into_iter()
|
||||
.map(|(message_id, result)| ProjectedMessage {
|
||||
message_id,
|
||||
role: Role::Tool,
|
||||
content: ProjectedContent::ToolResult(result),
|
||||
}),
|
||||
);
|
||||
Ok(Some((output, cursor)))
|
||||
}
|
||||
|
||||
fn project_message(message: &CanonicalMessage) -> ProjectedMessage {
|
||||
let content = match &message.content {
|
||||
MessageContent::Parts { parts } => ProjectedContent::Parts(parts.clone()),
|
||||
MessageContent::Assistant {
|
||||
text,
|
||||
thinking,
|
||||
replay_state,
|
||||
tool_calls,
|
||||
..
|
||||
} => ProjectedContent::Assistant {
|
||||
text: text.clone(),
|
||||
thinking: thinking.clone(),
|
||||
replay_state: replay_state.clone(),
|
||||
calls: tool_calls.clone(),
|
||||
},
|
||||
MessageContent::ToolResult(result) => ProjectedContent::ToolResult(result.clone()),
|
||||
};
|
||||
ProjectedMessage {
|
||||
message_id: message.message_id.clone(),
|
||||
role: message.role.clone(),
|
||||
content,
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,317 @@
|
||||
use std::{fmt, str::FromStr};
|
||||
|
||||
use reqwest::Url;
|
||||
use serde::{Deserialize, Serialize};
|
||||
use sha2::{Digest, Sha256};
|
||||
|
||||
use crate::{Error, Result};
|
||||
|
||||
#[derive(Clone, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)]
|
||||
pub enum ProviderType {
|
||||
#[serde(rename = "openai-chat")]
|
||||
OpenAiChat,
|
||||
#[serde(rename = "openai-responses")]
|
||||
OpenAiResponses,
|
||||
#[serde(rename = "anthropic")]
|
||||
Anthropic,
|
||||
}
|
||||
|
||||
impl ProviderType {
|
||||
pub fn as_str(self) -> &'static str {
|
||||
match self {
|
||||
Self::OpenAiChat => "openai-chat",
|
||||
Self::OpenAiResponses => "openai-responses",
|
||||
Self::Anthropic => "anthropic",
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
impl fmt::Display for ProviderType {
|
||||
fn fmt(&self, formatter: &mut fmt::Formatter<'_>) -> fmt::Result {
|
||||
formatter.write_str(self.as_str())
|
||||
}
|
||||
}
|
||||
|
||||
impl FromStr for ProviderType {
|
||||
type Err = Error;
|
||||
|
||||
fn from_str(value: &str) -> Result<Self> {
|
||||
match value {
|
||||
"openai-chat" => Ok(Self::OpenAiChat),
|
||||
"openai-responses" => Ok(Self::OpenAiResponses),
|
||||
"anthropic" => Ok(Self::Anthropic),
|
||||
_ => Err(Error::Config(format!("unsupported provider type: {value}"))),
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct ProviderEndpoint {
|
||||
pub provider_id: i64,
|
||||
pub name: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub base_url: String,
|
||||
pub has_api_key: bool,
|
||||
pub custom_headers: serde_json::Value,
|
||||
pub extra_params: serde_json::Value,
|
||||
pub created_at_ms: i64,
|
||||
pub updated_at_ms: i64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug)]
|
||||
pub struct ProviderEndpointSecret {
|
||||
pub endpoint: ProviderEndpoint,
|
||||
pub api_key: String,
|
||||
pub custom_headers: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize)]
|
||||
pub struct ProviderEndpointInput {
|
||||
pub name: String,
|
||||
pub provider_type: ProviderType,
|
||||
pub base_url: String,
|
||||
#[serde(default)]
|
||||
pub api_key: Option<String>,
|
||||
#[serde(default = "empty_object")]
|
||||
pub custom_headers: serde_json::Value,
|
||||
#[serde(default = "empty_object")]
|
||||
pub extra_params: serde_json::Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Deserialize, Serialize)]
|
||||
pub struct ProviderModelInput {
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub endpoint_type: ProviderType,
|
||||
#[serde(default)]
|
||||
pub request_url: String,
|
||||
#[serde(default = "enabled")]
|
||||
pub enabled: bool,
|
||||
#[serde(default)]
|
||||
pub sort_order: i64,
|
||||
pub context_window_tokens: Option<u64>,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
#[serde(default)]
|
||||
pub reasoning_enabled: bool,
|
||||
pub reasoning_effort: Option<String>,
|
||||
#[serde(default)]
|
||||
pub supports_image_generation: bool,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize)]
|
||||
pub struct ProviderModel {
|
||||
pub model_hash: String,
|
||||
pub provider_id: i64,
|
||||
pub model_id: String,
|
||||
pub display_name: String,
|
||||
pub endpoint_type: ProviderType,
|
||||
pub request_url: String,
|
||||
pub enabled: bool,
|
||||
pub sort_order: i64,
|
||||
pub context_window_tokens: Option<u64>,
|
||||
pub max_output_tokens: Option<u64>,
|
||||
pub reasoning_enabled: bool,
|
||||
pub reasoning_effort: Option<String>,
|
||||
pub supports_image_generation: bool,
|
||||
pub created_at_ms: i64,
|
||||
pub updated_at_ms: i64,
|
||||
}
|
||||
|
||||
impl ProviderModel {
|
||||
pub fn configure(&self, model: &mut super::ModelSpec) {
|
||||
model.display_name = Some(self.display_name.clone());
|
||||
model.supports_image_generation = self.supports_image_generation;
|
||||
model.reasoning.enabled |= self.reasoning_enabled;
|
||||
}
|
||||
}
|
||||
|
||||
pub fn normalize_base_url(value: &str) -> Result<String> {
|
||||
let mut url = Url::parse(value.trim())
|
||||
.map_err(|error| Error::Config(format!("invalid provider base URL: {error}")))?;
|
||||
if url.query().is_some() || url.fragment().is_some() {
|
||||
return Err(Error::Config(
|
||||
"provider base URL cannot contain query or fragment".into(),
|
||||
));
|
||||
}
|
||||
let path = url.path().trim_end_matches('/').to_string();
|
||||
url.set_path(if path.is_empty() { "/" } else { &path });
|
||||
Ok(url.as_str().trim_end_matches('/').to_string())
|
||||
}
|
||||
|
||||
pub fn model_hash(base_url: &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() {
|
||||
return Err(Error::Config("model id cannot be empty".into()));
|
||||
}
|
||||
let mut digest = Sha256::new();
|
||||
digest.update(base_url.as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(provider_type.as_str().as_bytes());
|
||||
digest.update([0]);
|
||||
digest.update(model_id.as_bytes());
|
||||
Ok(hex::encode(&digest.finalize()[..4]))
|
||||
}
|
||||
|
||||
pub fn resolve_request_url(
|
||||
base_url: &str,
|
||||
endpoint_type: ProviderType,
|
||||
request_url: &str,
|
||||
) -> Result<String> {
|
||||
let base_url = normalize_base_url(base_url)?;
|
||||
let request_url = request_url.trim();
|
||||
let combined = if request_url.starts_with("http://") || request_url.starts_with("https://") {
|
||||
let url = Url::parse(request_url)
|
||||
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
|
||||
if url.host_str().is_none() {
|
||||
return Err(Error::Config(
|
||||
"model request URL must contain a host".into(),
|
||||
));
|
||||
}
|
||||
url.to_string()
|
||||
} else {
|
||||
let path = if request_url.is_empty() {
|
||||
match endpoint_type {
|
||||
ProviderType::OpenAiChat => "/v1/chat/completions",
|
||||
ProviderType::OpenAiResponses => "/v1/responses",
|
||||
ProviderType::Anthropic => "/v1/messages",
|
||||
}
|
||||
} else if request_url.starts_with('/') {
|
||||
request_url
|
||||
} else {
|
||||
return Err(Error::Config(
|
||||
"model request URL must be an HTTP(S) URL or start with /".into(),
|
||||
));
|
||||
};
|
||||
format!("{}{}", base_url.trim_end_matches('/'), path)
|
||||
};
|
||||
let mut normalized = combined;
|
||||
while normalized.contains("/v1/v1") {
|
||||
normalized = normalized.replace("/v1/v1", "/v1");
|
||||
}
|
||||
Ok(normalized)
|
||||
}
|
||||
|
||||
pub fn is_sensitive_header(name: &str) -> bool {
|
||||
matches!(
|
||||
name.to_ascii_lowercase().as_str(),
|
||||
"authorization" | "proxy-authorization" | "x-api-key" | "api-key" | "cookie" | "set-cookie"
|
||||
)
|
||||
}
|
||||
|
||||
fn empty_object() -> serde_json::Value {
|
||||
serde_json::json!({})
|
||||
}
|
||||
|
||||
fn enabled() -> bool {
|
||||
true
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn hash_uses_normalized_url_type_and_model_only() {
|
||||
let first = model_hash(
|
||||
"HTTPS://Example.COM/v1/",
|
||||
ProviderType::OpenAiChat,
|
||||
"model-a",
|
||||
)
|
||||
.unwrap();
|
||||
let second = model_hash(
|
||||
"https://example.com/v1",
|
||||
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()
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_type_json_uses_public_identifiers() {
|
||||
for (value, provider_type) in [
|
||||
("openai-chat", ProviderType::OpenAiChat),
|
||||
("openai-responses", ProviderType::OpenAiResponses),
|
||||
("anthropic", ProviderType::Anthropic),
|
||||
] {
|
||||
assert_eq!(
|
||||
serde_json::from_str::<ProviderType>(&format!("\"{value}\"")).unwrap(),
|
||||
provider_type
|
||||
);
|
||||
assert_eq!(
|
||||
serde_json::to_string(&provider_type).unwrap(),
|
||||
format!("\"{value}\"")
|
||||
);
|
||||
}
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn resolves_default_relative_and_absolute_model_urls() {
|
||||
assert_eq!(
|
||||
resolve_request_url("https://example.com/v1", ProviderType::OpenAiChat, "").unwrap(),
|
||||
"https://example.com/v1/chat/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_request_url("https://example.com/v1", ProviderType::OpenAiResponses, "")
|
||||
.unwrap(),
|
||||
"https://example.com/v1/responses"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_request_url("https://example.com", ProviderType::Anthropic, "").unwrap(),
|
||||
"https://example.com/v1/messages"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_request_url(
|
||||
"https://example.com",
|
||||
ProviderType::OpenAiChat,
|
||||
"/v2/chat/completions"
|
||||
)
|
||||
.unwrap(),
|
||||
"https://example.com/v2/chat/completions"
|
||||
);
|
||||
assert_eq!(
|
||||
resolve_request_url(
|
||||
"https://example.com/v1",
|
||||
ProviderType::OpenAiChat,
|
||||
"https://gateway.example/v1/v1/custom"
|
||||
)
|
||||
.unwrap(),
|
||||
"https://gateway.example/v1/custom"
|
||||
);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn requested_runtime_limits_are_not_overridden_by_provider_config() {
|
||||
let provider = ProviderModel {
|
||||
model_hash: "12345678".into(),
|
||||
provider_id: 1,
|
||||
model_id: "model".into(),
|
||||
display_name: "Model".into(),
|
||||
endpoint_type: ProviderType::OpenAiResponses,
|
||||
request_url: String::new(),
|
||||
enabled: true,
|
||||
sort_order: 0,
|
||||
context_window_tokens: Some(200_000),
|
||||
max_output_tokens: None,
|
||||
reasoning_enabled: false,
|
||||
reasoning_effort: None,
|
||||
supports_image_generation: false,
|
||||
created_at_ms: 0,
|
||||
updated_at_ms: 0,
|
||||
};
|
||||
let mut selected = super::super::ModelSpec::new("12345678");
|
||||
selected.context_window_tokens = Some(800_000);
|
||||
provider.configure(&mut selected);
|
||||
assert_eq!(selected.context_window_tokens, Some(800_000));
|
||||
|
||||
let mut defaulted = super::super::ModelSpec::new("12345678");
|
||||
provider.configure(&mut defaulted);
|
||||
assert_eq!(defaulted.context_window_tokens, None);
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,58 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{
|
||||
CanonicalMessage, ConversationId, ModelSpec, PromptSpec, RevisionId, RunId, ToolCall,
|
||||
ToolRoundAssistant,
|
||||
};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub enum SubagentKind {
|
||||
GeneralPurpose,
|
||||
Named(String),
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub enum RunKind {
|
||||
Root,
|
||||
Subagent {
|
||||
parent_run_id: RunId,
|
||||
parent_tool_call_id: String,
|
||||
kind: SubagentKind,
|
||||
background: bool,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub enum SubagentModelOverride {
|
||||
Explicit(ModelSpec),
|
||||
Inherit,
|
||||
Disabled,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub enum RunAction {
|
||||
Start,
|
||||
Compact,
|
||||
Resume {
|
||||
pending_tool_round: Option<RecoveredToolRound>,
|
||||
},
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct RecoveredToolRound {
|
||||
pub assistant: ToolRoundAssistant,
|
||||
pub calls: Vec<ToolCall>,
|
||||
pub started_at_ms: u64,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct PreparedRun {
|
||||
pub run_id: RunId,
|
||||
pub conversation_id: ConversationId,
|
||||
pub kind: RunKind,
|
||||
pub model: ModelSpec,
|
||||
pub prompt: PromptSpec,
|
||||
pub initial_messages: Vec<CanonicalMessage>,
|
||||
pub action: RunAction,
|
||||
pub base_revision_id: RevisionId,
|
||||
}
|
||||
@@ -0,0 +1,23 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
use super::{CanonicalMessage, ContentPart, MessageContent, Origin, Role};
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct RuntimeEvent {
|
||||
pub event_id: String,
|
||||
pub text: String,
|
||||
}
|
||||
|
||||
impl RuntimeEvent {
|
||||
pub fn into_message(self) -> CanonicalMessage {
|
||||
CanonicalMessage {
|
||||
message_id: format!("runtime:{}", self.event_id),
|
||||
role: Role::User,
|
||||
origin: Origin::Runtime,
|
||||
content: MessageContent::Parts {
|
||||
parts: vec![ContentPart::Text { text: self.text }],
|
||||
},
|
||||
runtime_event_id: Some(self.event_id),
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,44 @@
|
||||
use serde::{Deserialize, Serialize};
|
||||
use serde_json::Value;
|
||||
|
||||
use super::ProviderReplayState;
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolDefinition {
|
||||
pub name: String,
|
||||
pub description: String,
|
||||
pub parameters: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolCall {
|
||||
pub index: usize,
|
||||
pub call_id: String,
|
||||
pub model_call_id: String,
|
||||
pub name: String,
|
||||
pub arguments_text: String,
|
||||
pub arguments: Value,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct ToolResult {
|
||||
pub call_id: String,
|
||||
pub content: String,
|
||||
pub is_error: bool,
|
||||
pub image: Option<ToolImageReference>,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq, Eq)]
|
||||
pub struct ToolImageReference {
|
||||
pub blob_id: String,
|
||||
pub mime_type: String,
|
||||
pub path: String,
|
||||
}
|
||||
|
||||
#[derive(Clone, Debug, Serialize, Deserialize, PartialEq)]
|
||||
pub struct ToolRoundAssistant {
|
||||
pub text: String,
|
||||
pub thinking: String,
|
||||
pub model_call_id: String,
|
||||
pub replay_state: Option<ProviderReplayState>,
|
||||
}
|
||||
@@ -0,0 +1,59 @@
|
||||
use std::ops::AddAssign;
|
||||
|
||||
use serde::{Deserialize, Serialize};
|
||||
|
||||
#[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 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;
|
||||
|
||||
#[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);
|
||||
}
|
||||
}
|
||||
Reference in New Issue
Block a user