merge old config

This commit is contained in:
leookun
2026-08-26 15:46:11 +08:00
parent 34f97334dc
commit e1937233ec
70 changed files with 4515 additions and 3764 deletions
+1
View File
@@ -39,6 +39,7 @@ rcgen = { version = "0.14", features = ["aws_lc_rs", "pem", "x509-parser"] }
scraper = "0.24"
serde = { version = "1", features = ["derive"] }
serde_json = "1"
serde_yaml = "0.9"
semble-core = { path = "../crates/semble-core" }
sha1 = "0.10"
sha2 = "0.10"
@@ -0,0 +1,197 @@
PRAGMA defer_foreign_keys = ON;
CREATE TABLE model_configs (
model_hash TEXT PRIMARY KEY,
sort_order INTEGER NOT NULL DEFAULT 0,
display_name TEXT NOT NULL,
model_type TEXT NOT NULL CHECK(model_type IN ('openai', 'anthropic')),
base_url TEXT NOT NULL,
use_full_url INTEGER NOT NULL DEFAULT 0 CHECK(use_full_url IN (0, 1)),
api_key TEXT NOT NULL,
tooltip_data TEXT NOT NULL,
model_id TEXT NOT NULL,
reasoning_effort TEXT,
openai_endpoint TEXT NOT NULL DEFAULT '',
openai_extra_params_enabled INTEGER NOT NULL DEFAULT 0 CHECK(openai_extra_params_enabled IN (0, 1)),
openai_extra_params_json TEXT NOT NULL DEFAULT '{}',
custom_headers_enabled INTEGER NOT NULL DEFAULT 0 CHECK(custom_headers_enabled IN (0, 1)),
custom_headers_json TEXT NOT NULL DEFAULT '{}',
anthropic_extra_params_enabled INTEGER NOT NULL DEFAULT 0 CHECK(anthropic_extra_params_enabled IN (0, 1)),
anthropic_extra_params_json TEXT NOT NULL DEFAULT '{}',
context_window_tokens INTEGER,
max_completion_tokens INTEGER,
anthropic_max_tokens INTEGER,
anthropic_thinking_effort TEXT,
thinking_budget_tokens INTEGER,
created_at_ms INTEGER NOT NULL,
updated_at_ms INTEGER NOT NULL
);
INSERT INTO model_configs (
model_hash,
sort_order,
display_name,
model_type,
base_url,
use_full_url,
api_key,
tooltip_data,
model_id,
reasoning_effort,
openai_endpoint,
openai_extra_params_enabled,
openai_extra_params_json,
custom_headers_enabled,
custom_headers_json,
anthropic_extra_params_enabled,
anthropic_extra_params_json,
context_window_tokens,
max_completion_tokens,
anthropic_max_tokens,
anthropic_thinking_effort,
thinking_budget_tokens,
created_at_ms,
updated_at_ms
)
SELECT
model.model_hash,
model.sort_order,
model.display_name,
CASE model.endpoint_type WHEN 'anthropic' THEN 'anthropic' ELSE 'openai' END,
CASE
WHEN model.request_url = '' THEN endpoint.base_url
WHEN model.request_url LIKE 'http://%' OR model.request_url LIKE 'https://%' THEN model.request_url
ELSE replace(rtrim(endpoint.base_url, '/') || '/' || ltrim(model.request_url, '/'), '/v1/v1/', '/v1/')
END,
CASE WHEN model.request_url = '' THEN 0 ELSE 1 END,
endpoint.api_key,
model.display_name,
model.model_id,
CASE
WHEN model.endpoint_type != 'anthropic' AND model.reasoning_enabled = 1
THEN COALESCE(NULLIF(trim(model.reasoning_effort), ''), 'medium')
ELSE NULL
END,
CASE model.endpoint_type
WHEN 'openai-responses' THEN '/v1/responses'
WHEN 'openai-chat' THEN '/v1/chat/completions'
ELSE ''
END,
CASE WHEN model.endpoint_type != 'anthropic' AND endpoint.extra_params_json != '{}' THEN 1 ELSE 0 END,
CASE WHEN model.endpoint_type != 'anthropic' THEN endpoint.extra_params_json ELSE '{}' END,
CASE WHEN endpoint.custom_headers_json != '{}' THEN 1 ELSE 0 END,
endpoint.custom_headers_json,
CASE WHEN model.endpoint_type = 'anthropic' AND endpoint.extra_params_json != '{}' THEN 1 ELSE 0 END,
CASE WHEN model.endpoint_type = 'anthropic' THEN endpoint.extra_params_json ELSE '{}' END,
model.context_window_tokens,
CASE WHEN model.endpoint_type != 'anthropic' THEN model.max_output_tokens ELSE NULL END,
CASE WHEN model.endpoint_type = 'anthropic' THEN model.max_output_tokens ELSE NULL END,
CASE WHEN model.endpoint_type = 'anthropic' THEN 'xhigh' ELSE NULL END,
NULL,
model.created_at_ms,
model.updated_at_ms
FROM provider_models AS model
JOIN provider_endpoints AS endpoint ON endpoint.provider_id = model.provider_id;
CREATE TABLE llm_calls_new (
call_id TEXT PRIMARY KEY,
run_id TEXT NOT NULL,
conversation_id TEXT NOT NULL,
provider_call_index INTEGER NOT NULL,
model_hash TEXT,
provider_type TEXT NOT NULL,
provider_url TEXT NOT NULL,
request_type TEXT NOT NULL,
request_url TEXT NOT NULL,
model_id TEXT NOT NULL,
display_name TEXT NOT NULL,
status TEXT NOT NULL,
finish_reason TEXT,
created_at_ms INTEGER NOT NULL,
request_started_at_ms INTEGER,
response_headers_at_ms INTEGER,
first_event_at_ms INTEGER,
first_text_at_ms INTEGER,
finished_at_ms INTEGER,
queue_ms INTEGER,
ttfb_ms INTEGER,
ttft_ms INTEGER,
duration_ms INTEGER,
input_tokens INTEGER,
output_tokens INTEGER,
total_tokens INTEGER,
cache_read_tokens INTEGER,
cache_write_tokens INTEGER,
reasoning_tokens INTEGER,
usage_json TEXT,
message_count INTEGER NOT NULL,
tool_count INTEGER NOT NULL,
request_bytes INTEGER,
response_bytes INTEGER NOT NULL DEFAULT 0,
stream_event_count INTEGER NOT NULL DEFAULT 0,
http_status INTEGER,
error_kind TEXT,
error_message TEXT,
detailed INTEGER NOT NULL,
reasoning_effort TEXT,
fast INTEGER NOT NULL DEFAULT 0 CHECK (fast IN (0, 1)),
FOREIGN KEY(model_hash) REFERENCES model_configs(model_hash)
);
INSERT INTO llm_calls_new (
call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type,
provider_url, request_type, request_url, model_id, display_name, status, finish_reason,
created_at_ms, request_started_at_ms, response_headers_at_ms, first_event_at_ms,
first_text_at_ms, finished_at_ms, queue_ms, ttfb_ms, ttft_ms, duration_ms,
input_tokens, output_tokens, total_tokens, cache_read_tokens, cache_write_tokens,
reasoning_tokens, usage_json, message_count, tool_count, request_bytes, response_bytes,
stream_event_count, http_status, error_kind, error_message, detailed, reasoning_effort, fast
)
SELECT
call_id, run_id, conversation_id, provider_call_index, model_hash, provider_type,
provider_url, request_type, request_url, model_id, display_name, status, finish_reason,
created_at_ms, request_started_at_ms, response_headers_at_ms, first_event_at_ms,
first_text_at_ms, finished_at_ms, queue_ms, ttfb_ms, ttft_ms, duration_ms,
input_tokens, output_tokens, total_tokens, cache_read_tokens, cache_write_tokens,
reasoning_tokens, usage_json, message_count, tool_count, request_bytes, response_bytes,
stream_event_count, http_status, error_kind, error_message, detailed, reasoning_effort, fast
FROM llm_calls;
CREATE TABLE llm_call_requests_new (
call_id TEXT PRIMARY KEY,
headers_json TEXT NOT NULL,
body_json TEXT NOT NULL,
byte_count INTEGER NOT NULL,
FOREIGN KEY(call_id) REFERENCES llm_calls_new(call_id) ON DELETE CASCADE
);
INSERT INTO llm_call_requests_new(call_id, headers_json, body_json, byte_count)
SELECT call_id, headers_json, body_json, byte_count FROM llm_call_requests;
CREATE TABLE llm_call_response_chunks_new (
call_id TEXT NOT NULL,
seq INTEGER NOT NULL,
received_offset_ms INTEGER NOT NULL,
data BLOB NOT NULL,
byte_count INTEGER NOT NULL,
PRIMARY KEY(call_id, seq),
FOREIGN KEY(call_id) REFERENCES llm_calls_new(call_id) ON DELETE CASCADE
);
INSERT INTO llm_call_response_chunks_new(call_id, seq, received_offset_ms, data, byte_count)
SELECT call_id, seq, received_offset_ms, data, byte_count FROM llm_call_response_chunks;
DROP TABLE llm_call_requests;
DROP TABLE llm_call_response_chunks;
DROP TABLE llm_calls;
DROP TABLE provider_models;
DROP TABLE provider_endpoints;
ALTER TABLE llm_calls_new RENAME TO llm_calls;
ALTER TABLE llm_call_requests_new RENAME TO llm_call_requests;
ALTER TABLE llm_call_response_chunks_new RENAME TO llm_call_response_chunks;
CREATE INDEX model_configs_sort ON model_configs(sort_order, display_name);
CREATE INDEX llm_calls_created ON llm_calls(created_at_ms DESC);
CREATE INDEX llm_calls_run ON llm_calls(run_id, provider_call_index);
CREATE INDEX llm_calls_model ON llm_calls(model_hash, created_at_ms DESC);
+10
View File
@@ -7,6 +7,8 @@ use crate::{Error, Result};
const DATA_DIR_NAME: &str = ".cursor-byok-v3";
const DATABASE_FILE_NAME: &str = "cursor-byok.db";
const V0049_DATA_DIR_NAME: &str = ".cursor-local-assistant-v2";
const V0049_CONFIG_FILE_NAME: &str = "config.yaml";
pub fn managed_data_dir() -> Result<PathBuf> {
let home_dir = dirs::home_dir()
@@ -18,6 +20,14 @@ pub fn managed_data_dir() -> Result<PathBuf> {
Ok(data_dir)
}
pub fn v0049_config_path() -> Result<PathBuf> {
let home_dir = dirs::home_dir()
.ok_or_else(|| Error::Config("cannot resolve user home directory".into()))?;
Ok(home_dir
.join(V0049_DATA_DIR_NAME)
.join(V0049_CONFIG_FILE_NAME))
}
#[derive(Clone, Debug, PartialEq, Eq)]
pub enum ProviderKind {
OpenAiChat,
-72
View File
@@ -1,72 +0,0 @@
use axum::{extract::State, http::StatusCode, Json};
use serde::{Deserialize, Serialize};
use crate::{
model::{ProviderEndpoint, ProviderEndpointInput, ProviderModel, ProviderModelInput},
Result,
};
use super::{ControlService, DiscoveredModels};
#[derive(Debug, Deserialize)]
#[serde(tag = "kind", rename_all = "snake_case")]
pub enum ProviderSelection {
Existing { provider_id: i64 },
New { input: ProviderEndpointInput },
}
#[derive(Debug, Deserialize)]
pub struct CreateCursorModels {
pub provider: ProviderSelection,
pub models: Vec<ProviderModelInput>,
}
#[derive(Debug, Serialize)]
pub struct CreatedCursorModels {
pub provider: ProviderEndpoint,
pub models: Vec<ProviderModel>,
}
#[derive(Debug, Deserialize)]
pub struct DiscoverCursorModels {
pub provider: ProviderSelection,
}
pub async fn create(
State(service): State<ControlService>,
Json(input): Json<CreateCursorModels>,
) -> Result<(StatusCode, Json<CreatedCursorModels>)> {
let (provider, models) = match input.provider {
ProviderSelection::Existing { provider_id } => {
let provider = service
.providers()
.await?
.into_iter()
.find(|provider| provider.provider_id == provider_id)
.ok_or_else(|| crate::Error::RunNotFound(format!("provider {provider_id}")))?;
let models = service.save_models(provider_id, &input.models).await?;
(provider, models)
}
ProviderSelection::New { input: provider } => {
service
.create_provider_with_models(&provider, &input.models)
.await?
}
};
Ok((
StatusCode::CREATED,
Json(CreatedCursorModels { provider, models }),
))
}
pub async fn discover(
State(service): State<ControlService>,
Json(input): Json<DiscoverCursorModels>,
) -> Result<Json<DiscoveredModels>> {
match input.provider {
ProviderSelection::Existing { provider_id } => {
Ok(Json(service.discover_models(provider_id).await?))
}
ProviderSelection::New { input } => Ok(Json(service.discover_input(&input).await?)),
}
}
+10 -27
View File
@@ -1,10 +1,8 @@
mod ads;
mod calls;
mod cursor_models;
mod harness;
mod models;
mod overview;
mod providers;
mod service;
mod settings;
@@ -22,8 +20,8 @@ use tower_http::{
use url::{Host, Url};
pub use service::{
CallDetail, CallSummary, ControlService, DiscoveredModels, ModelConnectivityResult,
ObservabilitySettings,
CallDetail, CallSummary, ControlService, DiscoveredModels, LegacyModelImportPreview,
LegacyModelImportResult, ModelConnectivityResult, ModelDiscoveryInput, ObservabilitySettings,
};
pub fn web_router(service: ControlService, assets: impl AsRef<std::path::Path>) -> Router {
@@ -117,22 +115,15 @@ pub fn api_router(service: ControlService) -> Router {
post(ads::dismiss),
)
.route(
"/__byok-api__/api/providers",
get(providers::list).post(providers::create),
"/__byok-api__/api/models",
get(models::list).post(models::create),
)
.route("/__byok-api__/api/models/discover", post(models::discover))
.route(
"/__byok-api__/api/providers/{provider_id}",
put(providers::update).delete(providers::remove),
"/__byok-api__/api/models/import-v0049",
get(models::preview_v0049).post(models::import_v0049),
)
.route(
"/__byok-api__/api/providers/{provider_id}/models/discover",
post(models::discover),
)
.route(
"/__byok-api__/api/providers/{provider_id}/models",
post(models::save),
)
.route("/__byok-api__/api/models", get(models::list))
.route("/__byok-api__/api/models/order", put(models::reorder))
.route("/__byok-api__/api/overview", get(overview::get))
.route(
"/__byok-api__/api/models/{model_hash}",
@@ -180,14 +171,6 @@ pub fn api_router(service: ControlService) -> Router {
"/__byok-api__/api/harness/cursor/enabled",
put(harness::set_enabled),
)
.route(
"/__byok-api__/api/harness/cursor/models",
post(cursor_models::create),
)
.route(
"/__byok-api__/api/harness/cursor/models/discover",
post(cursor_models::discover),
)
.with_state(service)
.layer(desktop_cors())
}
@@ -265,7 +248,7 @@ mod tests {
.clone()
.oneshot(
Request::builder()
.uri("/__byok-api__/api/providers")
.uri("/__byok-api__/api/models")
.header(header::ORIGIN, "tauri://localhost")
.body(Body::empty())
.unwrap(),
@@ -293,7 +276,7 @@ mod tests {
let response = router
.oneshot(
Request::builder()
.uri("/api/providers")
.uri("/api/models")
.body(Body::empty())
.unwrap(),
)
+38 -12
View File
@@ -6,32 +6,46 @@ use axum::{
use serde::Deserialize;
use crate::{
model::{ProviderModel, ProviderModelInput},
model::{ModelConfig, ModelConfigInput},
Result,
};
use super::{ControlService, DiscoveredModels, ModelConnectivityResult};
use super::{
ControlService, DiscoveredModels, LegacyModelImportPreview, LegacyModelImportResult,
ModelConnectivityResult, ModelDiscoveryInput,
};
#[derive(Deserialize)]
pub struct SaveModels {
pub models: Vec<ProviderModelInput>,
pub models: Vec<ModelConfigInput>,
}
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ProviderModel>>> {
#[derive(Deserialize)]
pub struct ModelOrder {
pub model_hashes: Vec<String>,
}
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ModelConfig>>> {
Ok(Json(service.models().await?))
}
pub async fn save(
pub async fn create(
State(service): State<ControlService>,
Path(provider_id): Path<i64>,
Json(input): Json<SaveModels>,
) -> Result<(StatusCode, Json<Vec<ProviderModel>>)> {
) -> Result<(StatusCode, Json<Vec<ModelConfig>>)> {
Ok((
StatusCode::CREATED,
Json(service.save_models(provider_id, &input.models).await?),
Json(service.create_models(&input.models).await?),
))
}
pub async fn reorder(
State(service): State<ControlService>,
Json(input): Json<ModelOrder>,
) -> Result<Json<Vec<ModelConfig>>> {
Ok(Json(service.reorder_models(&input.model_hashes).await?))
}
pub async fn remove(
State(service): State<ControlService>,
Path(model_hash): Path<String>,
@@ -43,8 +57,8 @@ pub async fn remove(
pub async fn update(
State(service): State<ControlService>,
Path(model_hash): Path<String>,
Json(input): Json<ProviderModelInput>,
) -> Result<Json<ProviderModel>> {
Json(input): Json<ModelConfigInput>,
) -> Result<Json<ModelConfig>> {
Ok(Json(service.update_model(&model_hash, &input).await?))
}
@@ -57,7 +71,19 @@ pub async fn test(
pub async fn discover(
State(service): State<ControlService>,
Path(provider_id): Path<i64>,
Json(input): Json<ModelDiscoveryInput>,
) -> Result<Json<DiscoveredModels>> {
Ok(Json(service.discover_models(provider_id).await?))
Ok(Json(service.discover_models(&input).await?))
}
pub async fn import_v0049(
State(service): State<ControlService>,
) -> Result<Json<LegacyModelImportResult>> {
Ok(Json(service.import_v0049_models().await?))
}
pub async fn preview_v0049(
State(service): State<ControlService>,
) -> Result<Json<LegacyModelImportPreview>> {
Ok(Json(service.preview_v0049_models().await?))
}
+1 -7
View File
@@ -15,7 +15,6 @@ pub struct OverviewRange {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<String>,
provider_ids: Option<String>,
}
pub async fn get(
@@ -24,12 +23,7 @@ pub async fn get(
) -> Result<Json<Overview>> {
Ok(Json(
service
.overview(
range.start_ms,
range.end_ms,
range.model_hashes.as_deref(),
range.provider_ids.as_deref(),
)
.overview(range.start_ms, range.end_ms, range.model_hashes.as_deref())
.await?,
))
}
-42
View File
@@ -1,42 +0,0 @@
use axum::{
extract::{Path, State},
http::StatusCode,
Json,
};
use crate::{
model::{ProviderEndpoint, ProviderEndpointInput},
Result,
};
use super::ControlService;
pub async fn list(State(service): State<ControlService>) -> Result<Json<Vec<ProviderEndpoint>>> {
Ok(Json(service.providers().await?))
}
pub async fn create(
State(service): State<ControlService>,
Json(input): Json<ProviderEndpointInput>,
) -> Result<(StatusCode, Json<ProviderEndpoint>)> {
Ok((
StatusCode::CREATED,
Json(service.create_provider(&input).await?),
))
}
pub async fn update(
State(service): State<ControlService>,
Path(provider_id): Path<i64>,
Json(input): Json<ProviderEndpointInput>,
) -> Result<Json<ProviderEndpoint>> {
Ok(Json(service.update_provider(provider_id, &input).await?))
}
pub async fn remove(
State(service): State<ControlService>,
Path(provider_id): Path<i64>,
) -> Result<StatusCode> {
service.delete_provider(provider_id).await?;
Ok(StatusCode::NO_CONTENT)
}
+147 -119
View File
@@ -16,9 +16,9 @@ use crate::{
harness::CursorHarness,
model::{
ContentPart, CursorRunTraceArtifact, CursorRunTraceSummary, LlmCallRequest,
LlmCallResponseChunk, LlmCallSummary, ModelInvocation, ModelRequest, ModelSpec, Overview,
ProjectedContent, ProjectedMessage, PromptSpec, ProviderEndpoint, ProviderEndpointInput,
ProviderModel, ProviderModelInput, ProviderType, Role,
LlmCallResponseChunk, LlmCallSummary, ModelConfig, ModelConfigInput, ModelInvocation,
ModelRequest, ModelSpec, ModelType, Overview, ProjectedContent, ProjectedMessage,
PromptSpec, ProviderType, Role,
},
provider::{ModelEvent, Provider},
store::{
@@ -40,6 +40,53 @@ pub struct DiscoveredModels {
pub models: Vec<String>,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportResult {
pub imported: usize,
pub skipped: usize,
pub total: usize,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportPreview {
pub source: String,
pub total: usize,
pub new_models: usize,
pub existing_models: usize,
pub models: Vec<LegacyModelImportPreviewItem>,
}
#[derive(Clone, Debug, Serialize)]
pub struct LegacyModelImportPreviewItem {
pub model_hash: String,
pub display_name: String,
pub model_id: String,
#[serde(rename = "type")]
pub model_type: ModelType,
pub existing: bool,
}
#[derive(Clone, Debug, Deserialize)]
pub struct ModelDiscoveryInput {
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
pub api_key: String,
#[serde(default)]
pub custom_headers_enabled: bool,
#[serde(default = "empty_json_object")]
pub custom_headers: serde_json::Value,
}
fn empty_json_object() -> serde_json::Value {
serde_json::json!({})
}
fn empty_json_object_ref() -> &'static serde_json::Value {
static EMPTY: std::sync::OnceLock<serde_json::Value> = std::sync::OnceLock::new();
EMPTY.get_or_init(empty_json_object)
}
#[derive(Clone, Debug, Serialize)]
pub struct ModelConnectivityResult {
pub duration_ms: u64,
@@ -163,28 +210,8 @@ impl ControlService {
Ok(())
}
pub async fn providers(&self) -> Result<Vec<ProviderEndpoint>> {
self.store.providers().await
}
pub async fn create_provider(&self, input: &ProviderEndpointInput) -> Result<ProviderEndpoint> {
self.store.create_provider(input).await
}
pub async fn update_provider(
&self,
provider_id: i64,
input: &ProviderEndpointInput,
) -> Result<ProviderEndpoint> {
self.store.update_provider(provider_id, input).await
}
pub async fn delete_provider(&self, provider_id: i64) -> Result<()> {
self.store.delete_provider(provider_id).await
}
pub async fn models(&self) -> Result<Vec<ProviderModel>> {
self.store.provider_models(false).await
pub async fn models(&self) -> Result<Vec<ModelConfig>> {
self.store.models().await
}
pub async fn overview(
@@ -192,31 +219,28 @@ impl ControlService {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<&str>,
provider_ids: Option<&str>,
) -> Result<Overview> {
self.store
.overview(start_ms, end_ms, model_hashes, provider_ids)
.await
self.store.overview(start_ms, end_ms, model_hashes).await
}
pub async fn save_models(
&self,
provider_id: i64,
models: &[ProviderModelInput],
) -> Result<Vec<ProviderModel>> {
self.store.save_provider_models(provider_id, models).await
pub async fn create_models(&self, models: &[ModelConfigInput]) -> Result<Vec<ModelConfig>> {
self.store.create_models(models).await
}
pub async fn reorder_models(&self, model_hashes: &[String]) -> Result<Vec<ModelConfig>> {
self.store.reorder_models(model_hashes).await
}
pub async fn delete_model(&self, model_hash: &str) -> Result<()> {
self.store.delete_provider_model(model_hash).await
self.store.delete_model(model_hash).await
}
pub async fn update_model(
&self,
model_hash: &str,
input: &ProviderModelInput,
) -> Result<ProviderModel> {
self.store.update_provider_model(model_hash, input).await
input: &ModelConfigInput,
) -> Result<ModelConfig> {
self.store.update_model(model_hash, input).await
}
pub async fn test_model(&self, model_hash: &str) -> Result<ModelConnectivityResult> {
@@ -225,19 +249,12 @@ impl ControlService {
let configured = self
.store
.provider_model(model_hash)
.model(model_hash)
.await?
.ok_or_else(|| Error::RunNotFound(format!("model {model_hash}")))?;
let mut model = ModelSpec::new(model_hash);
if configured.reasoning_enabled {
model.reasoning.enabled = true;
model.reasoning.effort = Some(
configured
.reasoning_effort
.filter(|effort| !effort.trim().is_empty())
.unwrap_or_else(|| "medium".into()),
);
}
configured.configure(&mut model);
model.max_output_tokens = Some(configured.max_output_tokens().unwrap_or(65_536));
let test_id = format!("model-test-{}", uuid::Uuid::new_v4());
let call_id = test_id.clone();
let invocation = ModelInvocation {
@@ -317,6 +334,11 @@ impl ControlService {
}
let elapsed = started.elapsed();
let output = output.trim().to_string();
if first_text_at.is_none() {
return Err(Error::Provider(
"model connectivity test received no text output".into(),
));
}
let tokens_estimated = output_tokens.is_none();
let output_tokens = output_tokens.unwrap_or_else(|| estimate_output_tokens(&output));
Ok(ModelConnectivityResult {
@@ -338,44 +360,58 @@ impl ControlService {
})
}
pub async fn create_provider_with_models(
&self,
provider: &ProviderEndpointInput,
models: &[ProviderModelInput],
) -> Result<(ProviderEndpoint, Vec<ProviderModel>)> {
self.store
.create_provider_with_models(provider, models)
.await
}
pub async fn discover_input(&self, input: &ProviderEndpointInput) -> Result<DiscoveredModels> {
pub async fn discover_models(&self, input: &ModelDiscoveryInput) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let base_url = crate::model::normalize_base_url(&input.base_url)?;
discover_provider_models(
let base_url = crate::model::normalize_request_url(&input.base_url)?;
discover_models_from_endpoint(
&client,
input.provider_type,
match input.model_type {
ModelType::OpenAi => ProviderType::OpenAiResponses,
ModelType::Anthropic => ProviderType::Anthropic,
},
&base_url,
input.api_key.as_deref().unwrap_or_default(),
&input.custom_headers,
&input.api_key,
if input.custom_headers_enabled {
&input.custom_headers
} else {
empty_json_object_ref()
},
)
.await
}
pub async fn discover_models(&self, provider_id: i64) -> Result<DiscoveredModels> {
let client = crate::network::client(&self.store).await?;
let provider = self
.store
.provider(provider_id)
.await?
.ok_or_else(|| Error::RunNotFound(format!("provider {provider_id}")))?;
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 import_v0049_models(&self) -> Result<LegacyModelImportResult> {
let path = crate::config::v0049_config_path()?;
let outcome = self.store.import_v0049_model_config(&path).await?;
Ok(LegacyModelImportResult {
imported: outcome.imported,
skipped: outcome.skipped,
total: outcome.total,
})
}
pub async fn preview_v0049_models(&self) -> Result<LegacyModelImportPreview> {
let path = crate::config::v0049_config_path()?;
let plan = self.store.preview_v0049_model_config(&path).await?;
let total = plan.models.len();
let existing_models = plan.models.iter().filter(|model| model.existing).count();
Ok(LegacyModelImportPreview {
source: path.display().to_string(),
total,
new_models: total - existing_models,
existing_models,
models: plan
.models
.into_iter()
.map(|model| LegacyModelImportPreviewItem {
model_hash: model.model_hash,
display_name: model.input.display_name,
model_id: model.input.model_id,
model_type: model.input.model_type,
existing: model.existing,
})
.collect(),
})
}
pub async fn calls(&self, limit: i64) -> Result<Vec<CallSummary>> {
@@ -595,7 +631,7 @@ fn readable_utf8(data: &[u8]) -> Option<&str> {
.then_some(value)
}
async fn discover_provider_models(
async fn discover_models_from_endpoint(
client: &reqwest::Client,
provider_type: ProviderType,
base_url: &str,
@@ -617,10 +653,10 @@ async fn discover_provider_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}")))?;
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
if url.host_str().is_none() {
return Err(Error::Config(
"provider base URL must contain a host".into(),
"model request URL must contain a host".into(),
));
}
url.set_path("/v1/models");
@@ -756,10 +792,7 @@ mod tests {
use tokio_util::sync::CancellationToken;
use crate::{
model::{
ModelInvocation, ProjectedContent, ProviderEndpointInput, ProviderModelInput,
ProviderType,
},
model::{ModelConfigInput, ModelInvocation, ModelType, ProjectedContent},
provider::{FinishReason, ModelEvent, Provider, ProviderStream},
store::Store,
};
@@ -803,34 +836,30 @@ mod tests {
.await
.unwrap();
let invocation = Arc::new(Mutex::new(None));
let provider = store
.create_provider(&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiResponses,
base_url: "https://example.com/v1".into(),
api_key: None,
custom_headers: serde_json::json!({}),
extra_params: serde_json::json!({}),
})
.await
.unwrap();
let model = store
.save_provider_model(
provider.provider_id,
&ProviderModelInput {
model_id: "reasoning-model".into(),
display_name: "Reasoning Model".into(),
endpoint_type: ProviderType::OpenAiResponses,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: true,
reasoning_effort: None,
supports_image_generation: false,
},
)
.create_model(&ModelConfigInput {
model_id: "reasoning-model".into(),
display_name: "Reasoning Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/responses".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Reasoning Model".into(),
sort_order: 0,
reasoning_effort: Some("medium".into()),
openai_endpoint: "/v1/responses".into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let service = ControlService::new(
@@ -925,16 +954,15 @@ mod tests {
)
.unwrap();
let result = service
.discover_input(&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiResponses,
.discover_models(&super::ModelDiscoveryInput {
model_type: ModelType::OpenAi,
base_url: format!("http://{address}/custom/responses"),
api_key: Some("secret".into()),
api_key: "secret".into(),
custom_headers_enabled: true,
custom_headers: serde_json::json!({
"uSeR-aGeNt": "inherited-user-agent",
"x-tenant": "tenant-a"
}),
extra_params: serde_json::json!({ "temperature": 0.7 }),
})
.await
.unwrap();
+1 -1
View File
@@ -128,7 +128,7 @@ async fn bidi_append_handler(
let conversation_id = decoded.conversation_id().map(str::to_owned);
let trace_metadata = decoded.trace_metadata();
let local = if let Some(model_id) = decoded.model_id() {
if registry.store().provider_model(model_id).await?.is_some() {
if registry.store().model(model_id).await?.is_some() {
tracing::info!(
request_id = decoded.request_id,
model_id,
+46 -56
View File
@@ -1,5 +1,3 @@
use std::collections::HashMap;
use axum::{
body::{Body, Bytes},
extract::{Extension, State},
@@ -14,7 +12,7 @@ use crate::{
proxy::{self, CursorProxy},
CursorSessionRegistry,
},
model::{format_token_count, parse_token_count, ProviderModel},
model::{format_token_count, parse_token_count, ModelConfig, ModelType},
Error, Result,
};
@@ -205,7 +203,7 @@ const EFFORTS: [(&str, &str); 5] = [
];
const DEFAULT_CONTEXT: &str = "200k";
fn context_options(model: &ProviderModel) -> Vec<(String, String)> {
fn context_options(model: &ModelConfig) -> Vec<(String, String)> {
let mut contexts = CONTEXTS
.into_iter()
.map(|(value, display_name)| (value.to_owned(), display_name.to_owned()))
@@ -227,30 +225,12 @@ pub async fn available_models(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().provider_models(true).await?;
let provider_names = registry
.store()
.providers()
.await?
.into_iter()
.map(|provider| (provider.provider_id, provider.name))
.collect::<HashMap<_, _>>();
let models = registry.store().models().await?;
tracing::info!(
model_count = models.len(),
"appending BYOK models to Cursor AvailableModels"
);
let available_models = models
.iter()
.map(|model| {
let provider_name = provider_names.get(&model.provider_id).ok_or_else(|| {
Error::Config(format!(
"provider {} for model {} does not exist",
model.provider_id, model.model_hash
))
})?;
Ok(available_model(model, provider_name))
})
.collect::<Result<Vec<_>>>()?;
let available_models = models.iter().map(available_model).collect::<Vec<_>>();
let local = AvailableModelsAddition {
model_names: models
.iter()
@@ -273,7 +253,7 @@ pub async fn usable_models(
Extension(proxy): Extension<CursorProxy>,
request: Request<Body>,
) -> Result<Response<Body>> {
let models = registry.store().provider_models(true).await?;
let models = registry.store().models().await?;
tracing::info!(
model_count = models.len(),
"appending BYOK models to Cursor GetUsableModels"
@@ -340,14 +320,14 @@ fn unary_payload(body: &Bytes) -> Result<(bool, &[u8])> {
Ok((true, &body[5..]))
}
fn available_model(model: &ProviderModel, provider_name: &str) -> AvailableModel {
fn available_model(model: &ModelConfig) -> AvailableModel {
let contexts = context_options(model);
let variants = model_variants(model, &contexts);
let legacy_slugs = variants
.iter()
.filter_map(|variant| variant.legacy_slug.clone())
.collect();
let tooltip = model_tooltip(model, "200K", "high", false);
let tooltip = model_tooltip(model);
AvailableModel {
name: model.model_hash.clone(),
default_on: true,
@@ -376,7 +356,10 @@ fn available_model(model: &ProviderModel, provider_name: &str) -> AvailableModel
display_name: "Cursor".into(),
}),
model_picker_badges: vec![ModelPickerBadge {
label: provider_name.into(),
label: match model.model_type {
ModelType::OpenAi => "OpenAI".into(),
ModelType::Anthropic => "Anthropic".into(),
},
variant: 1,
dismiss_on_selection: false,
}],
@@ -447,7 +430,7 @@ fn model_parameters(contexts: &[(String, String)]) -> Vec<ModelParameterDefiniti
]
}
fn model_variants(model: &ProviderModel, contexts: &[(String, String)]) -> Vec<ModelVariant> {
fn model_variants(model: &ModelConfig, contexts: &[(String, String)]) -> Vec<ModelVariant> {
let mut variants = Vec::with_capacity(contexts.len() * EFFORTS.len() * 2);
for (context, context_name) in contexts {
for (effort, effort_name) in EFFORTS {
@@ -467,7 +450,7 @@ fn model_variants(model: &ProviderModel, contexts: &[(String, String)]) -> Vec<M
}
fn model_variant(
model: &ProviderModel,
model: &ModelConfig,
context: &str,
context_name: &str,
effort: &str,
@@ -507,7 +490,7 @@ fn model_variant(
is_max_mode: false,
is_default_max_config: is_default.then_some(true),
is_default_non_max_config: is_default.then_some(true),
tooltip_data: Some(model_tooltip(model, context_name, effort, fast)),
tooltip_data: Some(model_tooltip(model)),
display_name_outside_picker: Some(display_name),
variant_string_representation: Some(format!(
"{}[context={context},effort={effort},fast={fast}]",
@@ -521,22 +504,13 @@ fn model_variant(
}
}
fn model_tooltip(
model: &ProviderModel,
context_name: &str,
effort: &str,
fast: bool,
) -> TooltipData {
let fast_label = if fast { " (Fast)" } else { "" };
fn model_tooltip(model: &ModelConfig) -> TooltipData {
TooltipData {
markdown_content: Some(format!(
"**{}{fast_label}**<br /><br />{context_name} context window<br /><br />*Version: {effort} effort*",
model.display_name
)),
markdown_content: Some(model.tooltip_data.clone()),
}
}
fn usable_model(model: &ProviderModel) -> agent::ModelDetails {
fn usable_model(model: &ModelConfig) -> agent::ModelDetails {
agent::ModelDetails {
model_id: model.model_hash.clone(),
display_model_id: model.model_hash.clone(),
@@ -555,25 +529,34 @@ mod tests {
#[test]
fn maps_byok_model_to_cursor_catalog_fields() {
let model = ProviderModel {
let model = ModelConfig {
model_hash: "33ceed20".into(),
provider_id: 1,
model_id: "deepseek-v4-flash".into(),
display_name: "DeepSeek V4 Flash".into(),
endpoint_type: crate::model::ProviderType::OpenAiResponses,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: Some(272_000),
max_output_tokens: None,
reasoning_enabled: false,
display_name: "DeepSeek V4 Flash".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/responses".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "DeepSeek V4 Flash".into(),
model_id: "deepseek-v4-flash".into(),
reasoning_effort: None,
supports_image_generation: false,
openai_endpoint: "/v1/responses".into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: Some(272_000),
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
created_at_ms: 0,
updated_at_ms: 0,
};
let mapped = available_model(&model, "OpenRouter");
let mapped = available_model(&model);
assert_eq!(mapped.name, "33ceed20");
assert!(mapped.default_on);
assert_eq!(mapped.supports_agent, Some(true));
@@ -591,6 +574,13 @@ mod tests {
);
assert_eq!(mapped.server_model_name.as_deref(), Some("33ceed20"));
assert_eq!(mapped.named_model_section_index, Some(1));
assert_eq!(
mapped
.tooltip_data
.as_ref()
.and_then(|tooltip| tooltip.markdown_content.as_deref()),
Some("DeepSeek V4 Flash")
);
assert_eq!(mapped.vendor_name.as_deref(), Some("cursor"));
assert_eq!(mapped.parameter_definitions.len(), 3);
let context = mapped
@@ -643,7 +633,7 @@ mod tests {
assert_eq!(mapped.variants.len(), 50);
assert_eq!(mapped.legacy_slugs.len(), 50);
assert_eq!(mapped.model_picker_badges.len(), 1);
assert_eq!(mapped.model_picker_badges[0].label, "OpenRouter");
assert_eq!(mapped.model_picker_badges[0].label, "OpenAI");
assert!(!mapped.model_picker_badges[0].dismiss_on_selection);
let default = mapped
.variants
+2 -6
View File
@@ -135,12 +135,8 @@ pub(crate) async fn prepare(
mode_from_proto(mode_number)?
};
let mut model = model::requested_model(request)?;
if let Some(provider_model) = store
.provider_model(&model.model_id)
.await?
.filter(|model| model.enabled)
{
provider_model.configure(&mut model);
if let Some(configured_model) = store.model(&model.model_id).await? {
configured_model.configure(&mut model);
}
let dynamic = context::dynamic_mcp(request, &request_context)?;
let subagent_model_overrides = model::overrides(request)?;
+2 -2
View File
@@ -94,9 +94,9 @@ impl CursorHarness {
}
pub async fn status(&self) -> Result<CursorHarnessStatus> {
let models = self.inner.store.provider_models(false).await?;
let models = self.inner.store.models().await?;
let configured_models = models.len();
let enabled_models = models.iter().filter(|model| model.enabled).count();
let enabled_models = configured_models;
let ca = self.inner.ca.state()?;
if integration_prerequisites_ready(&ca, self.inner.backend_addr.read().is_some()) {
self.enable().await?;
+508
View File
@@ -0,0 +1,508 @@
use std::{fmt, str::FromStr};
use reqwest::Url;
use serde::{Deserialize, Serialize};
use sha2::{Digest, Sha256};
use crate::{Error, Result};
pub const OPENAI_RESPONSES_ENDPOINT: &str = "/v1/responses";
pub const OPENAI_CHAT_ENDPOINT: &str = "/v1/chat/completions";
#[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, Copy, Debug, Deserialize, Serialize, PartialEq, Eq)]
#[serde(rename_all = "lowercase")]
pub enum ModelType {
OpenAi,
Anthropic,
}
impl ModelType {
pub fn as_str(self) -> &'static str {
match self {
Self::OpenAi => "openai",
Self::Anthropic => "anthropic",
}
}
}
impl FromStr for ModelType {
type Err = Error;
fn from_str(value: &str) -> Result<Self> {
match value {
"openai" => Ok(Self::OpenAi),
"anthropic" => Ok(Self::Anthropic),
_ => Err(Error::Config(format!("unsupported model type: {value}"))),
}
}
}
#[derive(Clone, Debug, Deserialize, Serialize)]
pub struct ModelConfigInput {
#[serde(default)]
pub sort_order: i64,
pub display_name: String,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
#[serde(default)]
pub use_full_url: bool,
pub api_key: String,
pub tooltip_data: String,
pub model_id: String,
#[serde(default)]
pub reasoning_effort: Option<String>,
#[serde(default)]
pub openai_endpoint: String,
#[serde(default)]
pub openai_extra_params_enabled: bool,
#[serde(default = "empty_object")]
pub openai_extra_params: serde_json::Value,
#[serde(default)]
pub custom_headers_enabled: bool,
#[serde(default = "empty_object")]
pub custom_headers: serde_json::Value,
#[serde(default)]
pub anthropic_extra_params_enabled: bool,
#[serde(default = "empty_object")]
pub anthropic_extra_params: serde_json::Value,
pub context_window_tokens: Option<u64>,
pub max_completion_tokens: Option<u64>,
pub anthropic_max_tokens: Option<u64>,
#[serde(default)]
pub anthropic_thinking_effort: Option<String>,
pub thinking_budget_tokens: Option<u64>,
}
#[derive(Clone, Debug, Serialize)]
pub struct ModelConfig {
pub model_hash: String,
pub sort_order: i64,
pub display_name: String,
#[serde(rename = "type")]
pub model_type: ModelType,
pub base_url: String,
pub use_full_url: bool,
pub api_key: String,
pub tooltip_data: String,
pub model_id: String,
pub reasoning_effort: Option<String>,
pub openai_endpoint: String,
pub openai_extra_params_enabled: bool,
pub openai_extra_params: serde_json::Value,
pub custom_headers_enabled: bool,
pub custom_headers: serde_json::Value,
pub anthropic_extra_params_enabled: bool,
pub anthropic_extra_params: serde_json::Value,
pub context_window_tokens: Option<u64>,
pub max_completion_tokens: Option<u64>,
pub anthropic_max_tokens: Option<u64>,
pub anthropic_thinking_effort: Option<String>,
pub thinking_budget_tokens: Option<u64>,
pub created_at_ms: i64,
pub updated_at_ms: i64,
}
impl ModelConfig {
pub fn provider_type(&self) -> ProviderType {
match self.model_type {
ModelType::Anthropic => ProviderType::Anthropic,
ModelType::OpenAi if self.openai_endpoint == OPENAI_RESPONSES_ENDPOINT => {
ProviderType::OpenAiResponses
}
ModelType::OpenAi => ProviderType::OpenAiChat,
}
}
pub fn request_url(&self) -> Result<String> {
resolve_request_url(
self.model_type,
&self.base_url,
&self.openai_endpoint,
self.use_full_url,
)
}
pub fn max_output_tokens(&self) -> Option<u64> {
match self.model_type {
ModelType::OpenAi => self.max_completion_tokens,
ModelType::Anthropic => self.anthropic_max_tokens.or(self.max_completion_tokens),
}
}
pub fn extra_params(&self) -> &serde_json::Value {
match self.model_type {
ModelType::OpenAi if self.openai_extra_params_enabled => &self.openai_extra_params,
ModelType::Anthropic if self.anthropic_extra_params_enabled => {
&self.anthropic_extra_params
}
_ => empty_object_ref(),
}
}
pub fn configure(&self, model: &mut super::ModelSpec) {
model.display_name = Some(self.display_name.clone());
if model.reasoning.effort.is_none() {
model.reasoning.effort = match self.model_type {
ModelType::OpenAi => self.reasoning_effort.clone(),
ModelType::Anthropic => self.anthropic_thinking_effort.clone(),
};
}
model.reasoning.enabled |= model.reasoning.effort.is_some();
}
}
pub fn normalize_model_input(input: &ModelConfigInput) -> Result<ModelConfigInput> {
let display_name = required(&input.display_name, "model display name")?;
let base_url = normalize_request_url(&input.base_url)?;
let api_key = required(&input.api_key, "model API key")?;
let tooltip_data = required(&input.tooltip_data, "model tooltip")?;
let model_id = required(&input.model_id, "model id")?;
let reasoning_effort = normalize_effort(input.reasoning_effort.as_deref(), true)?;
let anthropic_thinking_effort = match input.model_type {
ModelType::Anthropic => Some(
normalize_effort(
input.anthropic_thinking_effort.as_deref().or(Some("xhigh")),
false,
)?
.expect("Anthropic effort has a default"),
),
ModelType::OpenAi => None,
};
let openai_endpoint = match input.model_type {
ModelType::OpenAi => normalize_openai_endpoint(&input.openai_endpoint)?,
ModelType::Anthropic => String::new(),
};
validate_object(&input.openai_extra_params, "OpenAI extra params")?;
validate_object(&input.anthropic_extra_params, "Anthropic extra params")?;
validate_headers(&input.custom_headers)?;
let normalized = ModelConfigInput {
sort_order: input.sort_order.max(0),
display_name,
model_type: input.model_type,
base_url,
use_full_url: input.use_full_url,
api_key,
tooltip_data,
model_id,
reasoning_effort: (input.model_type == ModelType::OpenAi)
.then_some(reasoning_effort)
.flatten(),
openai_endpoint,
openai_extra_params_enabled: input.model_type == ModelType::OpenAi
&& input.openai_extra_params_enabled,
openai_extra_params: if input.model_type == ModelType::OpenAi {
input.openai_extra_params.clone()
} else {
empty_object()
},
custom_headers_enabled: input.custom_headers_enabled,
custom_headers: input.custom_headers.clone(),
anthropic_extra_params_enabled: input.model_type == ModelType::Anthropic
&& input.anthropic_extra_params_enabled,
anthropic_extra_params: if input.model_type == ModelType::Anthropic {
input.anthropic_extra_params.clone()
} else {
empty_object()
},
context_window_tokens: positive(input.context_window_tokens, "context window")?,
max_completion_tokens: positive(input.max_completion_tokens, "max completion tokens")?,
anthropic_max_tokens: positive(input.anthropic_max_tokens, "Anthropic max tokens")?,
anthropic_thinking_effort,
thinking_budget_tokens: positive(input.thinking_budget_tokens, "thinking budget")?,
};
resolve_request_url(
normalized.model_type,
&normalized.base_url,
&normalized.openai_endpoint,
normalized.use_full_url,
)?;
Ok(normalized)
}
pub fn model_hash(input: &ModelConfigInput) -> Result<String> {
let normalized = normalize_model_input(input)?;
let request_url = resolve_request_url(
normalized.model_type,
&normalized.base_url,
&normalized.openai_endpoint,
normalized.use_full_url,
)?;
let mut parts = vec![
request_url,
normalized.model_id,
normalized.api_key,
normalized.display_name,
];
if normalized.model_type == ModelType::OpenAi {
parts.push(normalized.openai_endpoint);
}
let digest = Sha256::digest(parts.join("\n").as_bytes());
Ok(hex::encode(&digest[..8]))
}
pub fn normalize_request_url(value: &str) -> Result<String> {
let value = value.trim();
let url = Url::parse(value)
.map_err(|error| Error::Config(format!("invalid model request URL: {error}")))?;
if !matches!(url.scheme(), "http" | "https") || url.host_str().is_none() {
return Err(Error::Config(
"model request URL must be an HTTP(S) URL with a host".into(),
));
}
if url.fragment().is_some() {
return Err(Error::Config(
"model request URL cannot contain a fragment".into(),
));
}
Ok(value.into())
}
pub fn resolve_request_url(
model_type: ModelType,
base_url: &str,
openai_endpoint: &str,
use_full_url: bool,
) -> Result<String> {
let base_url = normalize_request_url(base_url)?;
let endpoint = match model_type {
ModelType::OpenAi => normalize_openai_endpoint(openai_endpoint)?,
ModelType::Anthropic => "/v1/messages".into(),
};
if use_full_url {
return Ok(base_url);
}
append_standard_endpoint(&base_url, &endpoint)
}
fn append_standard_endpoint(base_url: &str, endpoint: &str) -> Result<String> {
let mut url = Url::parse(base_url)
.map_err(|error| Error::Config(format!("invalid model server URL: {error}")))?;
let base_path = url.path().trim_end_matches('/').to_string();
let endpoint = if has_trailing_version(&base_path) {
endpoint.strip_prefix("/v1").unwrap_or(endpoint)
} else {
endpoint
};
url.set_path(&format!("{base_path}{endpoint}"));
normalize_request_url(url.as_str())
}
fn has_trailing_version(path: &str) -> bool {
let Some(segment) = path.rsplit('/').next() else {
return false;
};
segment.strip_prefix('v').is_some_and(|digits| {
!digits.is_empty() && digits.bytes().all(|byte| byte.is_ascii_digit())
})
}
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 normalize_openai_endpoint(value: &str) -> Result<String> {
match value.trim() {
"" | OPENAI_RESPONSES_ENDPOINT => Ok(OPENAI_RESPONSES_ENDPOINT.into()),
OPENAI_CHAT_ENDPOINT => Ok(OPENAI_CHAT_ENDPOINT.into()),
value => Err(Error::Config(format!(
"unsupported OpenAI endpoint: {value}"
))),
}
}
fn normalize_effort(value: Option<&str>, allow_empty: bool) -> Result<Option<String>> {
let value = value.unwrap_or_default().trim().to_ascii_lowercase();
if value.is_empty() && allow_empty {
return Ok(None);
}
if matches!(value.as_str(), "low" | "medium" | "high" | "xhigh" | "max") {
Ok(Some(value))
} else {
Err(Error::Config(format!(
"unsupported reasoning effort: {value}"
)))
}
}
fn positive(value: Option<u64>, label: &str) -> Result<Option<u64>> {
match value {
Some(0) => Err(Error::Config(format!("{label} must be greater than zero"))),
value => Ok(value),
}
}
fn required(value: &str, label: &str) -> Result<String> {
let value = value.trim();
if value.is_empty() {
Err(Error::Config(format!("{label} cannot be empty")))
} else {
Ok(value.into())
}
}
fn validate_object(value: &serde_json::Value, label: &str) -> Result<()> {
if value.is_object() {
Ok(())
} else {
Err(Error::Config(format!("{label} must be a JSON object")))
}
}
fn validate_headers(value: &serde_json::Value) -> Result<()> {
validate_object(value, "custom headers")?;
for (name, value) in value.as_object().expect("validated object") {
if name.trim().is_empty() || !value.is_string() {
return Err(Error::Config(
"custom headers must have non-empty names and string values".into(),
));
}
}
Ok(())
}
fn empty_object() -> serde_json::Value {
serde_json::json!({})
}
fn empty_object_ref() -> &'static serde_json::Value {
static EMPTY: std::sync::OnceLock<serde_json::Value> = std::sync::OnceLock::new();
EMPTY.get_or_init(empty_object)
}
#[cfg(test)]
mod tests {
use super::*;
fn input() -> ModelConfigInput {
ModelConfigInput {
sort_order: 1,
display_name: "Model A".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/custom/generate".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Model A".into(),
model_id: "model-a".into(),
reasoning_effort: Some("high".into()),
openai_endpoint: OPENAI_RESPONSES_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: empty_object(),
custom_headers_enabled: false,
custom_headers: empty_object(),
anthropic_extra_params_enabled: false,
anthropic_extra_params: empty_object(),
context_window_tokens: Some(200_000),
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
}
}
#[test]
fn hash_matches_the_v0049_channel_identity() {
let input = input();
let expected = Sha256::digest(
"https://example.com/custom/generate\nmodel-a\nsecret\nModel A\n/v1/responses"
.as_bytes(),
);
assert_eq!(model_hash(&input).unwrap(), hex::encode(&expected[..8]));
}
#[test]
fn request_url_is_exact_and_protocol_does_not_depend_on_its_path() {
assert_eq!(
resolve_request_url(
ModelType::OpenAi,
"https://example.com/custom/generate?api-version=2026-01-01",
OPENAI_RESPONSES_ENDPOINT,
true,
)
.unwrap(),
"https://example.com/custom/generate?api-version=2026-01-01"
);
assert_eq!(
resolve_request_url(
ModelType::OpenAi,
"https://example.com/another/arbitrary/path",
OPENAI_CHAT_ENDPOINT,
true,
)
.unwrap(),
"https://example.com/another/arbitrary/path"
);
assert_eq!(
resolve_request_url(ModelType::Anthropic, "https://example.com/claude", "", true)
.unwrap(),
"https://example.com/claude"
);
assert_eq!(
resolve_request_url(
ModelType::Anthropic,
"https://example.com/claude/",
"",
true
)
.unwrap(),
"https://example.com/claude/"
);
assert_eq!(
resolve_request_url(
ModelType::OpenAi,
"https://example.com/v1",
OPENAI_RESPONSES_ENDPOINT,
false,
)
.unwrap(),
"https://example.com/v1/responses"
);
assert_eq!(
resolve_request_url(ModelType::Anthropic, "https://example.com/v1", "", false).unwrap(),
"https://example.com/v1/messages"
);
}
}
+2 -2
View File
@@ -1,3 +1,4 @@
mod configuration;
mod conversation;
mod cursor_trace;
mod inference;
@@ -6,13 +7,13 @@ mod message;
mod model_spec;
mod overview;
mod projection;
mod provider;
mod run;
mod runtime_tag;
mod token_count;
mod tool;
mod usage;
pub use configuration::*;
pub use conversation::*;
pub use cursor_trace::*;
pub use inference::*;
@@ -21,7 +22,6 @@ pub use message::*;
pub use model_spec::*;
pub use overview::*;
pub use projection::*;
pub use provider::*;
pub use run::*;
pub use runtime_tag::*;
pub(crate) use token_count::*;
-341
View File
@@ -1,341 +0,0 @@
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 api_key: Option<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 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,
api_key: &str,
provider_type: ProviderType,
model_id: &str,
) -> Result<String> {
let base_url = normalize_base_url(base_url)?;
let model_id = model_id.trim();
if model_id.is_empty() {
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(api_key.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_key_type_and_model() {
let first = model_hash(
"HTTPS://Example.COM/v1/",
"secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap();
let second = model_hash(
"https://example.com/v1",
"secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap();
assert_eq!(first, second);
assert_ne!(
first,
model_hash(
"https://example.com/v1",
"different-secret",
ProviderType::OpenAiChat,
"model-a",
)
.unwrap()
);
assert_ne!(
first,
model_hash(
"https://example.com/v1",
"secret",
ProviderType::Anthropic,
"model-a",
)
.unwrap()
);
}
#[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);
}
}
+17 -21
View File
@@ -6,7 +6,7 @@ use tokio_util::sync::CancellationToken;
use crate::{
config::{ProviderConfig, ProviderKind},
model::{resolve_request_url, ModelInvocation, ModelLatency, NewLlmCall, ProviderType},
model::{ModelInvocation, ModelLatency, NewLlmCall, ProviderType},
store::Store,
Error, Result,
};
@@ -41,21 +41,13 @@ impl Provider for ProviderRouter {
Box::pin(try_stream! {
let selected = invocation.request.model.model_id.clone();
let model = store
.provider_model(&selected)
.model(&selected)
.await?
.filter(|model| model.enabled)
.ok_or_else(|| Error::Provider(format!("unknown or disabled model: {selected}")))?;
let endpoint = store
.provider(model.provider_id)
.await?
.ok_or_else(|| Error::Provider(format!("provider {} no longer exists", model.provider_id)))?;
let request_url = resolve_request_url(
&endpoint.endpoint.base_url,
model.endpoint_type,
&model.request_url,
)?;
.ok_or_else(|| Error::Provider(format!("unknown model: {selected}")))?;
let provider_type = model.provider_type();
let request_url = model.request_url()?;
model.configure(&mut invocation.request.model);
invocation.request.model.extra_params = endpoint.endpoint.extra_params.clone();
invocation.request.model.extra_params = model.extra_params().clone();
invocation.request.model.model_id = model.model_id.clone();
let recorder = CallRecorder::start(store.clone(), NewLlmCall {
call_id: invocation.call_id.clone(),
@@ -63,9 +55,9 @@ impl Provider for ProviderRouter {
conversation_id: invocation.conversation_id.clone(),
provider_call_index: invocation.provider_call_index.min(i64::MAX as u64) as i64,
model_hash: model.model_hash.clone(),
provider_type: endpoint.endpoint.provider_type,
provider_url: endpoint.endpoint.base_url.clone(),
request_type: model.endpoint_type,
provider_type,
provider_url: model.base_url.clone(),
request_type: provider_type,
request_url: request_url.clone(),
model_id: model.model_id.clone(),
display_name: model.display_name.clone(),
@@ -76,15 +68,19 @@ impl Provider for ProviderRouter {
detailed: false,
}).await?;
let config = ProviderConfig {
kind: match model.endpoint_type {
kind: match provider_type {
ProviderType::OpenAiChat => ProviderKind::OpenAiChat,
ProviderType::OpenAiResponses => ProviderKind::OpenAiResponses,
ProviderType::Anthropic => ProviderKind::Anthropic,
},
request_url,
api_key: endpoint.endpoint.api_key.clone().unwrap_or_default(),
custom_headers: custom_headers(&endpoint.custom_headers)?,
max_output_tokens: model.max_output_tokens,
api_key: model.api_key.clone(),
custom_headers: if model.custom_headers_enabled {
custom_headers(&model.custom_headers)?
} else {
reqwest::header::HeaderMap::new()
},
max_output_tokens: model.max_output_tokens(),
request_timeout,
};
let client = crate::network::client_builder(&store)
+390
View File
@@ -0,0 +1,390 @@
use std::{collections::HashSet, path::Path};
use serde::Deserialize;
use crate::{
model::{
model_hash, normalize_model_input, normalize_request_url, ModelConfigInput, ModelType,
OPENAI_CHAT_ENDPOINT, OPENAI_RESPONSES_ENDPOINT,
},
Error, Result,
};
use super::Store;
pub struct LegacyModelImportPlan {
pub models: Vec<LegacyModelImportEntry>,
}
pub struct LegacyModelImportEntry {
pub model_hash: String,
pub input: ModelConfigInput,
pub existing: bool,
}
pub struct LegacyModelImportOutcome {
pub imported: usize,
pub skipped: usize,
pub total: usize,
}
#[derive(Default, Deserialize)]
struct LegacyConfig {
#[serde(rename = "modelAdapters", default)]
model_adapters: Vec<LegacyModel>,
}
#[derive(Default, Deserialize)]
struct LegacyModel {
#[serde(default)]
sort: i64,
#[serde(rename = "displayName", default)]
display_name: String,
#[serde(rename = "type", default)]
model_type: String,
#[serde(rename = "baseURL", default)]
base_url: String,
#[serde(rename = "apiKey", default)]
api_key: String,
#[serde(rename = "tooltipData", default)]
tooltip_data: String,
#[serde(rename = "modelID", default)]
model_id: String,
#[serde(rename = "reasoningEffort", default)]
reasoning_effort: String,
#[serde(rename = "openAIEndpoint", default)]
openai_endpoint: String,
#[serde(rename = "openAIExtraParamsEnabled", default)]
openai_extra_params_enabled: bool,
#[serde(rename = "openAIExtraParamsJSON", default)]
openai_extra_params_json: String,
#[serde(rename = "customHeadersEnabled", default)]
custom_headers_enabled: bool,
#[serde(rename = "customHeadersJSON", default)]
custom_headers_json: String,
#[serde(rename = "anthropicExtraParamsEnabled", default)]
anthropic_extra_params_enabled: bool,
#[serde(rename = "anthropicExtraParamsJSON", default)]
anthropic_extra_params_json: String,
#[serde(rename = "contextWindowTokens", default)]
context_window_tokens: u64,
#[serde(rename = "maxCompletionTokens", default)]
max_completion_tokens: u64,
#[serde(rename = "anthropicMaxTokens", default)]
anthropic_max_tokens: u64,
#[serde(rename = "anthropicThinkingEffort", default)]
anthropic_thinking_effort: String,
#[serde(rename = "thinkingBudgetTokens", default)]
thinking_budget_tokens: u64,
}
impl Store {
pub async fn preview_v0049_model_config(&self, path: &Path) -> Result<LegacyModelImportPlan> {
let inputs = load_v0049_model_config(path)?;
let existing = self
.models()
.await?
.into_iter()
.map(|model| model.model_hash)
.collect::<HashSet<_>>();
let mut seen = HashSet::with_capacity(inputs.len());
let mut models = Vec::with_capacity(inputs.len());
for input in inputs {
let input = normalize_model_input(&input)?;
let hash = model_hash(&input)?;
if seen.insert(hash.clone()) {
models.push(LegacyModelImportEntry {
existing: existing.contains(&hash),
model_hash: hash,
input,
});
}
}
Ok(LegacyModelImportPlan { models })
}
pub async fn import_v0049_model_config(&self, path: &Path) -> Result<LegacyModelImportOutcome> {
let plan = self.preview_v0049_model_config(path).await?;
let total = plan.models.len();
let missing = plan
.models
.into_iter()
.filter(|model| !model.existing)
.map(|model| model.input)
.collect::<Vec<_>>();
let imported = self.create_models_if_missing(&missing).await?;
Ok(LegacyModelImportOutcome {
imported,
skipped: total - imported,
total,
})
}
}
fn load_v0049_model_config(path: &Path) -> Result<Vec<ModelConfigInput>> {
let raw = match std::fs::read(path) {
Ok(raw) => raw,
Err(error) if error.kind() == std::io::ErrorKind::NotFound => {
return Err(Error::Config(format!(
"v0.0.49 config not found at {}",
path.display()
)))
}
Err(error) => return Err(error.into()),
};
let legacy: LegacyConfig = serde_yaml::from_slice(&raw)
.map_err(|error| Error::Config(format!("invalid v0.0.49 config: {error}")))?;
if legacy.model_adapters.is_empty() {
return Err(Error::Config(
"v0.0.49 config contains no model adapters".into(),
));
}
let models = legacy
.model_adapters
.into_iter()
.map(model_input)
.collect::<Result<Vec<_>>>()?;
Ok(models)
}
fn model_input(model: LegacyModel) -> Result<ModelConfigInput> {
let model_type = match model.model_type.trim().to_ascii_lowercase().as_str() {
"openai" => ModelType::OpenAi,
"anthropic" => ModelType::Anthropic,
value => {
return Err(Error::Config(format!(
"unsupported v0.0.49 model type: {value}"
)))
}
};
let (base_url, openai_endpoint, use_full_url) =
legacy_request_configuration(model_type, &model.base_url, &model.openai_endpoint)?;
Ok(ModelConfigInput {
sort_order: model.sort,
display_name: model.display_name.clone(),
model_type,
base_url,
use_full_url,
api_key: model.api_key,
tooltip_data: if model.tooltip_data.trim().is_empty() {
model.display_name
} else {
model.tooltip_data
},
model_id: model.model_id,
reasoning_effort: optional_string(model.reasoning_effort),
openai_endpoint,
openai_extra_params_enabled: model.openai_extra_params_enabled,
openai_extra_params: enabled_json_object(
model_type == ModelType::OpenAi && model.openai_extra_params_enabled,
&model.openai_extra_params_json,
)?,
custom_headers_enabled: model.custom_headers_enabled,
custom_headers: enabled_json_object(
model.custom_headers_enabled,
&model.custom_headers_json,
)?,
anthropic_extra_params_enabled: model.anthropic_extra_params_enabled,
anthropic_extra_params: enabled_json_object(
model_type == ModelType::Anthropic && model.anthropic_extra_params_enabled,
&model.anthropic_extra_params_json,
)?,
context_window_tokens: positive(model.context_window_tokens),
max_completion_tokens: positive(model.max_completion_tokens),
anthropic_max_tokens: positive(model.anthropic_max_tokens),
anthropic_thinking_effort: optional_string(model.anthropic_thinking_effort),
thinking_budget_tokens: positive(model.thinking_budget_tokens),
})
}
fn legacy_request_configuration(
model_type: ModelType,
base_url: &str,
openai_endpoint: &str,
) -> Result<(String, String, bool)> {
let base_url = normalize_request_url(base_url)?;
match model_type {
ModelType::Anthropic => {
let use_full_url = url_path_ends_with(&base_url, "/messages");
Ok((base_url, String::new(), use_full_url))
}
ModelType::OpenAi => {
let detected = openai_protocol_from_url(&base_url);
let configured = match openai_endpoint.trim() {
"" | OPENAI_RESPONSES_ENDPOINT => OPENAI_RESPONSES_ENDPOINT,
OPENAI_CHAT_ENDPOINT => OPENAI_CHAT_ENDPOINT,
"/custom" => OPENAI_CHAT_ENDPOINT,
value => {
return Err(Error::Config(format!(
"unsupported v0.0.49 OpenAI endpoint: {value}"
)))
}
};
let protocol = detected.unwrap_or(configured);
let use_full_url = detected.is_some() || openai_endpoint.trim() == "/custom";
Ok((base_url, protocol.into(), use_full_url))
}
}
}
fn openai_protocol_from_url(value: &str) -> Option<&'static str> {
let url = reqwest::Url::parse(value).ok()?;
let path = url.path().trim_end_matches('/');
if path.to_ascii_lowercase().ends_with("/responses") {
Some(OPENAI_RESPONSES_ENDPOINT)
} else if path.to_ascii_lowercase().ends_with("/chat/completions") {
Some(OPENAI_CHAT_ENDPOINT)
} else {
None
}
}
fn url_path_ends_with(value: &str, suffix: &str) -> bool {
reqwest::Url::parse(value).is_ok_and(|url| {
url.path()
.trim_end_matches('/')
.to_ascii_lowercase()
.ends_with(suffix)
})
}
fn enabled_json_object(enabled: bool, value: &str) -> Result<serde_json::Value> {
if enabled {
json_object(value)
} else {
Ok(serde_json::json!({}))
}
}
fn json_object(value: &str) -> Result<serde_json::Value> {
if value.trim().is_empty() {
return Ok(serde_json::json!({}));
}
let value: serde_json::Value = serde_json::from_str(value)?;
if value.is_object() {
Ok(value)
} else {
Err(Error::Config(
"v0.0.49 model JSON fields must be objects".into(),
))
}
}
fn positive(value: u64) -> Option<u64> {
(value > 0).then_some(value)
}
fn optional_string(value: String) -> Option<String> {
(!value.trim().is_empty()).then_some(value)
}
#[cfg(test)]
mod tests {
use super::*;
#[tokio::test]
async fn manually_imports_v0049_models_as_complete_request_urls() {
let directory = tempfile::tempdir().unwrap();
let config = directory.path().join("config.yaml");
std::fs::write(
&config,
r#"modelAdapters:
- sort: 1
displayName: Model A
type: openai
baseURL: https://example.com/v1
apiKey: secret
tooltipData: Example model
modelID: model-a
reasoningEffort: high
openAIEndpoint: /v1/responses
openAIExtraParamsEnabled: true
openAIExtraParamsJSON: '{"service_tier":"priority"}'
customHeadersEnabled: true
customHeadersJSON: '{"x-client":"cursor-byok"}'
contextWindowTokens: 200000
maxCompletionTokens: 8192
- sort: 2
displayName: Custom Chat
type: openai
baseURL: https://example.com/proxy/generate?api-version=2026-01-01
apiKey: secret
modelID: model-b
openAIEndpoint: /custom
openAIExtraParamsEnabled: false
openAIExtraParamsJSON: not-valid-json
- sort: 3
displayName: Claude
type: anthropic
baseURL: https://example.com/anthropic
apiKey: secret
modelID: model-c
customHeadersEnabled: false
customHeadersJSON: not-valid-json
"#,
)
.unwrap();
let store = Store::connect("sqlite::memory:").await.unwrap();
let preview = store.preview_v0049_model_config(&config).await.unwrap();
assert_eq!(preview.models.len(), 3);
assert!(preview.models.iter().all(|model| !model.existing));
store.create_model(&preview.models[0].input).await.unwrap();
let preview = store.preview_v0049_model_config(&config).await.unwrap();
assert_eq!(
preview.models.iter().filter(|model| model.existing).count(),
1
);
let first = store.import_v0049_model_config(&config).await.unwrap();
assert_eq!(first.imported, 2);
assert_eq!(first.skipped, 1);
assert_eq!(first.total, 3);
let models = store.models().await.unwrap();
assert_eq!(models.len(), 3);
assert_eq!(models[0].model_hash.len(), 16);
assert_eq!(models[0].base_url, "https://example.com/v1");
assert!(!models[0].use_full_url);
assert_eq!(
models[0].request_url().unwrap(),
"https://example.com/v1/responses"
);
assert_eq!(models[0].openai_extra_params["service_tier"], "priority");
assert_eq!(
models[1].base_url,
"https://example.com/proxy/generate?api-version=2026-01-01"
);
assert_eq!(models[1].openai_endpoint, OPENAI_CHAT_ENDPOINT);
assert!(models[1].use_full_url);
assert_eq!(models[1].openai_extra_params, serde_json::json!({}));
assert_eq!(
models[2].request_url().unwrap(),
"https://example.com/anthropic/v1/messages"
);
assert!(!models[2].use_full_url);
assert_eq!(models[2].custom_headers, serde_json::json!({}));
let preview = store.preview_v0049_model_config(&config).await.unwrap();
assert!(preview.models.iter().all(|model| model.existing));
let second = store.import_v0049_model_config(&config).await.unwrap();
assert_eq!(second.imported, 0);
assert_eq!(second.skipped, 3);
assert_eq!(second.total, 3);
assert_eq!(store.models().await.unwrap().len(), 3);
}
#[test]
fn v0049_anthropic_full_request_url_is_not_modified() {
let (request_url, endpoint, use_full_url) = legacy_request_configuration(
ModelType::Anthropic,
"https://example.com/proxy/messages?api-version=2026-01-01",
"",
)
.unwrap();
assert_eq!(
request_url,
"https://example.com/proxy/messages?api-version=2026-01-01"
);
assert!(endpoint.is_empty());
assert!(use_full_url);
}
}
+29 -29
View File
@@ -328,36 +328,36 @@ fn summary_from_row(row: sqlx::sqlite::SqliteRow) -> Result<LlmCallSummary> {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{ProviderEndpointInput, ProviderModelInput};
use crate::model::{ModelConfigInput, ModelType};
#[tokio::test]
async fn latest_usage_anchor_uses_the_latest_completed_call_for_the_same_conversation_and_model(
) {
let store = Store::connect("sqlite::memory:").await.unwrap();
let (provider, model) = store
.create_provider_with_model(
&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiResponses,
base_url: "https://example.com".into(),
api_key: Some("secret".into()),
custom_headers: serde_json::json!({}),
extra_params: serde_json::json!({}),
},
&ProviderModelInput {
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: Some(16_000),
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: false,
},
)
let model = store
.create_model(&ModelConfigInput {
model_id: "model".into(),
display_name: "Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/responses".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Model".into(),
sort_order: 0,
reasoning_effort: None,
openai_endpoint: "/v1/responses".into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: Some(200_000),
max_completion_tokens: Some(16_000),
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let conversation_id = ConversationId::new("conversation");
@@ -374,10 +374,10 @@ mod tests {
conversation_id: conversation_id.to_string(),
provider_call_index: 0,
model_hash: model.model_hash.clone(),
provider_type: provider.provider_type,
provider_url: provider.base_url.clone(),
request_type: model.endpoint_type,
request_url: model.request_url.clone(),
provider_type: model.provider_type(),
provider_url: model.base_url.clone(),
request_type: model.provider_type(),
request_url: model.request_url().unwrap(),
model_id: model.model_id.clone(),
display_name: model.display_name.clone(),
reasoning_effort: None,
+2 -1
View File
@@ -2,10 +2,11 @@ mod cas;
mod conversations;
mod cursor_traces;
mod input_anchors;
mod legacy_config;
mod llm_calls;
mod messages;
mod models;
mod overview;
mod providers;
mod revisions;
mod runs;
mod settings;
+416
View File
@@ -0,0 +1,416 @@
use std::{collections::HashSet, str::FromStr};
use sqlx::{Row, Sqlite, Transaction};
use crate::{
model::{model_hash, normalize_model_input, ModelConfig, ModelConfigInput, ModelType},
Error, Result,
};
use super::{now_ms, Store};
const MODEL_COLUMNS: &str = r#"
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
max_completion_tokens, anthropic_max_tokens, anthropic_thinking_effort,
thinking_budget_tokens, created_at_ms, updated_at_ms
"#;
impl Store {
pub async fn models(&self) -> Result<Vec<ModelConfig>> {
let query =
format!("SELECT {MODEL_COLUMNS} FROM model_configs ORDER BY sort_order, display_name");
sqlx::query(&query)
.fetch_all(&self.pool)
.await?
.into_iter()
.map(model_from_row)
.collect()
}
pub async fn model(&self, hash: &str) -> Result<Option<ModelConfig>> {
let query = format!("SELECT {MODEL_COLUMNS} FROM model_configs WHERE model_hash = ?");
sqlx::query(&query)
.bind(hash)
.fetch_optional(&self.pool)
.await?
.map(model_from_row)
.transpose()
}
pub async fn create_model(&self, input: &ModelConfigInput) -> Result<ModelConfig> {
let mut models = self.create_models(std::slice::from_ref(input)).await?;
Ok(models.remove(0))
}
pub async fn create_models(&self, inputs: &[ModelConfigInput]) -> Result<Vec<ModelConfig>> {
if inputs.is_empty() {
return Err(Error::Config("at least one model is required".into()));
}
let mut normalized = Vec::with_capacity(inputs.len());
let mut hashes = HashSet::with_capacity(inputs.len());
for input in inputs {
let input = normalize_model_input(input)?;
let hash = model_hash(&input)?;
if !hashes.insert(hash.clone()) {
return Err(Error::Config("model configurations must be unique".into()));
}
normalized.push((hash, input));
}
let now = now_ms();
let mut transaction = self.pool.begin().await?;
for (hash, input) in &normalized {
insert_model(&mut transaction, hash, input, now).await?;
}
transaction.commit().await?;
let mut saved = Vec::with_capacity(normalized.len());
for (hash, _) in normalized {
saved.push(self.model(&hash).await?.expect("inserted model must exist"));
}
Ok(saved)
}
pub(super) async fn create_models_if_missing(
&self,
inputs: &[ModelConfigInput],
) -> Result<usize> {
let mut normalized = Vec::with_capacity(inputs.len());
let mut hashes = HashSet::with_capacity(inputs.len());
for input in inputs {
let input = normalize_model_input(input)?;
let hash = model_hash(&input)?;
if hashes.insert(hash.clone()) {
normalized.push((hash, input));
}
}
let now = now_ms();
let mut transaction = self.pool.begin().await?;
let mut inserted = 0;
for (hash, input) in &normalized {
inserted += usize::from(
insert_model_with_conflict(&mut transaction, hash, input, now, true).await?,
);
}
transaction.commit().await?;
Ok(inserted)
}
pub async fn update_model(
&self,
current_hash: &str,
input: &ModelConfigInput,
) -> Result<ModelConfig> {
let current = self
.model(current_hash)
.await?
.ok_or_else(|| Error::RunNotFound(format!("model {current_hash}")))?;
let input = normalize_model_input(input)?;
let next_hash = model_hash(&input)?;
let now = now_ms();
let mut transaction = self.pool.begin().await?;
if next_hash != current.model_hash {
sqlx::query("UPDATE llm_calls SET model_hash = NULL WHERE model_hash = ?")
.bind(&current.model_hash)
.execute(&mut *transaction)
.await?;
}
let result = sqlx::query(
r#"UPDATE model_configs SET
model_hash = ?, sort_order = ?, display_name = ?, model_type = ?, base_url = ?,
use_full_url = ?, api_key = ?, tooltip_data = ?, model_id = ?, reasoning_effort = ?,
openai_endpoint = ?, openai_extra_params_enabled = ?, openai_extra_params_json = ?,
custom_headers_enabled = ?, custom_headers_json = ?,
anthropic_extra_params_enabled = ?, anthropic_extra_params_json = ?,
context_window_tokens = ?, max_completion_tokens = ?, anthropic_max_tokens = ?,
anthropic_thinking_effort = ?, thinking_budget_tokens = ?, updated_at_ms = ?
WHERE model_hash = ?"#,
)
.bind(&next_hash)
.bind(input.sort_order)
.bind(&input.display_name)
.bind(input.model_type.as_str())
.bind(&input.base_url)
.bind(input.use_full_url)
.bind(&input.api_key)
.bind(&input.tooltip_data)
.bind(&input.model_id)
.bind(&input.reasoning_effort)
.bind(&input.openai_endpoint)
.bind(input.openai_extra_params_enabled)
.bind(serde_json::to_string(&input.openai_extra_params)?)
.bind(input.custom_headers_enabled)
.bind(serde_json::to_string(&input.custom_headers)?)
.bind(input.anthropic_extra_params_enabled)
.bind(serde_json::to_string(&input.anthropic_extra_params)?)
.bind(input.context_window_tokens.map(to_i64).transpose()?)
.bind(input.max_completion_tokens.map(to_i64).transpose()?)
.bind(input.anthropic_max_tokens.map(to_i64).transpose()?)
.bind(&input.anthropic_thinking_effort)
.bind(input.thinking_budget_tokens.map(to_i64).transpose()?)
.bind(now)
.bind(current_hash)
.execute(&mut *transaction)
.await?;
if result.rows_affected() != 1 {
return Err(Error::RunNotFound(format!("model {current_hash}")));
}
transaction.commit().await?;
Ok(self
.model(&next_hash)
.await?
.expect("updated model must exist"))
}
pub async fn delete_model(&self, hash: &str) -> Result<()> {
let mut transaction = self.pool.begin().await?;
sqlx::query("UPDATE llm_calls SET model_hash = NULL WHERE model_hash = ?")
.bind(hash)
.execute(&mut *transaction)
.await?;
let result = sqlx::query("DELETE FROM model_configs WHERE model_hash = ?")
.bind(hash)
.execute(&mut *transaction)
.await?;
if result.rows_affected() != 1 {
return Err(Error::RunNotFound(format!("model {hash}")));
}
transaction.commit().await?;
Ok(())
}
pub async fn reorder_models(&self, model_hashes: &[String]) -> Result<Vec<ModelConfig>> {
let current = self.models().await?;
let current_hashes = current
.iter()
.map(|model| model.model_hash.as_str())
.collect::<HashSet<_>>();
let requested_hashes = model_hashes
.iter()
.map(String::as_str)
.collect::<HashSet<_>>();
if model_hashes.len() != current.len()
|| requested_hashes.len() != current.len()
|| requested_hashes != current_hashes
{
return Err(Error::Config(
"model configuration changed; refresh and try sorting again".into(),
));
}
let now = now_ms();
let mut transaction = self.pool.begin().await?;
for (index, hash) in model_hashes.iter().enumerate() {
sqlx::query(
"UPDATE model_configs SET sort_order = ?, updated_at_ms = ? WHERE model_hash = ?",
)
.bind(i64::try_from(index + 1).expect("model order fits in i64"))
.bind(now)
.bind(hash)
.execute(&mut *transaction)
.await?;
}
transaction.commit().await?;
self.models().await
}
}
async fn insert_model(
transaction: &mut Transaction<'_, Sqlite>,
hash: &str,
input: &ModelConfigInput,
now: i64,
) -> Result<()> {
insert_model_with_conflict(transaction, hash, input, now, false).await?;
Ok(())
}
async fn insert_model_with_conflict(
transaction: &mut Transaction<'_, Sqlite>,
hash: &str,
input: &ModelConfigInput,
now: i64,
ignore_existing: bool,
) -> Result<bool> {
let mut statement = String::from(
r#"INSERT INTO model_configs(
model_hash, sort_order, display_name, model_type, base_url, use_full_url, api_key, tooltip_data,
model_id, reasoning_effort, openai_endpoint, openai_extra_params_enabled,
openai_extra_params_json, custom_headers_enabled, custom_headers_json,
anthropic_extra_params_enabled, anthropic_extra_params_json, context_window_tokens,
max_completion_tokens, anthropic_max_tokens, anthropic_thinking_effort,
thinking_budget_tokens, created_at_ms, updated_at_ms
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)"#,
);
if ignore_existing {
statement.push_str(" ON CONFLICT(model_hash) DO NOTHING");
}
let result = sqlx::query(&statement)
.bind(hash)
.bind(input.sort_order)
.bind(&input.display_name)
.bind(input.model_type.as_str())
.bind(&input.base_url)
.bind(input.use_full_url)
.bind(&input.api_key)
.bind(&input.tooltip_data)
.bind(&input.model_id)
.bind(&input.reasoning_effort)
.bind(&input.openai_endpoint)
.bind(input.openai_extra_params_enabled)
.bind(serde_json::to_string(&input.openai_extra_params)?)
.bind(input.custom_headers_enabled)
.bind(serde_json::to_string(&input.custom_headers)?)
.bind(input.anthropic_extra_params_enabled)
.bind(serde_json::to_string(&input.anthropic_extra_params)?)
.bind(input.context_window_tokens.map(to_i64).transpose()?)
.bind(input.max_completion_tokens.map(to_i64).transpose()?)
.bind(input.anthropic_max_tokens.map(to_i64).transpose()?)
.bind(&input.anthropic_thinking_effort)
.bind(input.thinking_budget_tokens.map(to_i64).transpose()?)
.bind(now)
.bind(now)
.execute(&mut **transaction)
.await?;
Ok(result.rows_affected() == 1)
}
fn model_from_row(row: sqlx::sqlite::SqliteRow) -> Result<ModelConfig> {
Ok(ModelConfig {
model_hash: row.try_get("model_hash")?,
sort_order: row.try_get("sort_order")?,
display_name: row.try_get("display_name")?,
model_type: ModelType::from_str(row.try_get("model_type")?)?,
base_url: row.try_get("base_url")?,
use_full_url: row.try_get("use_full_url")?,
api_key: row.try_get("api_key")?,
tooltip_data: row.try_get("tooltip_data")?,
model_id: row.try_get("model_id")?,
reasoning_effort: row.try_get("reasoning_effort")?,
openai_endpoint: row.try_get("openai_endpoint")?,
openai_extra_params_enabled: row.try_get("openai_extra_params_enabled")?,
openai_extra_params: serde_json::from_str(
row.try_get::<String, _>("openai_extra_params_json")?
.as_str(),
)?,
custom_headers_enabled: row.try_get("custom_headers_enabled")?,
custom_headers: serde_json::from_str(
row.try_get::<String, _>("custom_headers_json")?.as_str(),
)?,
anthropic_extra_params_enabled: row.try_get("anthropic_extra_params_enabled")?,
anthropic_extra_params: serde_json::from_str(
row.try_get::<String, _>("anthropic_extra_params_json")?
.as_str(),
)?,
context_window_tokens: optional_u64(&row, "context_window_tokens")?,
max_completion_tokens: optional_u64(&row, "max_completion_tokens")?,
anthropic_max_tokens: optional_u64(&row, "anthropic_max_tokens")?,
anthropic_thinking_effort: row.try_get("anthropic_thinking_effort")?,
thinking_budget_tokens: optional_u64(&row, "thinking_budget_tokens")?,
created_at_ms: row.try_get("created_at_ms")?,
updated_at_ms: row.try_get("updated_at_ms")?,
})
}
fn optional_u64(row: &sqlx::sqlite::SqliteRow, column: &str) -> Result<Option<u64>> {
row.try_get::<Option<i64>, _>(column)?
.map(|value| {
u64::try_from(value).map_err(|_| Error::Config(format!("{column} cannot be negative")))
})
.transpose()
}
fn to_i64(value: u64) -> Result<i64> {
i64::try_from(value).map_err(|_| Error::Config("token value is too large".into()))
}
#[cfg(test)]
mod tests {
use super::*;
fn input(name: &str) -> ModelConfigInput {
ModelConfigInput {
sort_order: 1,
display_name: name.into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/responses".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Example model".into(),
model_id: "model-a".into(),
reasoning_effort: Some("high".into()),
openai_endpoint: "/v1/responses".into(),
openai_extra_params_enabled: true,
openai_extra_params: serde_json::json!({"service_tier":"priority"}),
custom_headers_enabled: true,
custom_headers: serde_json::json!({"x-client":"cursor-byok"}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: Some(200_000),
max_completion_tokens: Some(8_192),
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
}
}
#[tokio::test]
async fn model_configuration_round_trips_and_updates_identity() {
let store = Store::connect("sqlite::memory:").await.unwrap();
let created = store.create_model(&input("Model A")).await.unwrap();
assert_eq!(created.model_hash.len(), 16);
assert_eq!(created.custom_headers["x-client"], "cursor-byok");
assert_eq!(store.models().await.unwrap().len(), 1);
let updated = store
.update_model(&created.model_hash, &input("Renamed"))
.await
.unwrap();
assert_ne!(updated.model_hash, created.model_hash);
assert!(store.model(&created.model_hash).await.unwrap().is_none());
store.delete_model(&updated.model_hash).await.unwrap();
assert!(store.models().await.unwrap().is_empty());
}
#[tokio::test]
async fn batch_creation_is_atomic() {
let store = Store::connect("sqlite::memory:").await.unwrap();
let duplicate = input("Model A");
assert!(store
.create_models(&[duplicate.clone(), duplicate])
.await
.is_err());
assert!(store.models().await.unwrap().is_empty());
}
#[tokio::test]
async fn model_order_is_replaced_atomically() {
let store = Store::connect("sqlite::memory:").await.unwrap();
let first = store.create_model(&input("First")).await.unwrap();
let mut second_input = input("Second");
second_input.model_id = "model-b".into();
second_input.sort_order = 2;
let second = store.create_model(&second_input).await.unwrap();
let reordered = store
.reorder_models(&[second.model_hash.clone(), first.model_hash.clone()])
.await
.unwrap();
assert_eq!(reordered[0].model_hash, second.model_hash);
assert_eq!(reordered[0].sort_order, 1);
assert_eq!(reordered[1].model_hash, first.model_hash);
assert_eq!(reordered[1].sort_order, 2);
assert!(store
.reorder_models(std::slice::from_ref(&first.model_hash))
.await
.is_err());
assert_eq!(
store.models().await.unwrap()[0].model_hash,
second.model_hash
);
}
}
+5 -25
View File
@@ -24,7 +24,6 @@ impl Store {
start_ms: Option<i64>,
end_ms: Option<i64>,
model_hashes: Option<&str>,
provider_ids: Option<&str>,
) -> Result<Overview> {
let call_row = sqlx::query(
"SELECT
@@ -35,11 +34,7 @@ impl Store {
WHERE status != 'running'
AND (? IS NULL OR created_at_ms >= ?)
AND (? IS NULL OR created_at_ms < ?)
AND (? IS NULL OR model_hash IN (SELECT value FROM json_each(?)))
AND (? IS NULL OR model_hash IN (
SELECT model_hash FROM provider_models
WHERE provider_id IN (SELECT value FROM json_each(?))
))",
AND (? IS NULL OR model_hash IN (SELECT value FROM json_each(?)))",
)
.bind(start_ms)
.bind(start_ms)
@@ -47,8 +42,6 @@ impl Store {
.bind(end_ms)
.bind(model_hashes)
.bind(model_hashes)
.bind(provider_ids)
.bind(provider_ids)
.fetch_one(&self.pool)
.await?;
let token_row = sqlx::query(&format!(
@@ -60,11 +53,7 @@ impl Store {
FROM llm_calls
WHERE (? IS NULL OR created_at_ms >= ?)
AND (? IS NULL OR created_at_ms < ?)
AND (? IS NULL OR model_hash IN (SELECT value FROM json_each(?)))
AND (? IS NULL OR model_hash IN (
SELECT model_hash FROM provider_models
WHERE provider_id IN (SELECT value FROM json_each(?))
))",
AND (? IS NULL OR model_hash IN (SELECT value FROM json_each(?)))",
fresh_input = fresh_input_sql(),
))
.bind(start_ms)
@@ -73,8 +62,6 @@ impl Store {
.bind(end_ms)
.bind(model_hashes)
.bind(model_hashes)
.bind(provider_ids)
.bind(provider_ids)
.fetch_one(&self.pool)
.await?;
@@ -108,10 +95,6 @@ impl Store {
WHERE created_at_ms >= ?
AND (? IS NULL OR created_at_ms < ?)
AND (? IS NULL OR model_hash IN (SELECT value FROM json_each(?)))
AND (? IS NULL OR model_hash IN (
SELECT model_hash FROM provider_models
WHERE provider_id IN (SELECT value FROM json_each(?))
))
GROUP BY bucket_start_ms
ORDER BY bucket_start_ms",
fresh_input = fresh_input_sql(),
@@ -121,8 +104,6 @@ impl Store {
.bind(end_ms)
.bind(model_hashes)
.bind(model_hashes)
.bind(provider_ids)
.bind(provider_ids)
.fetch_all(&self.pool)
.await?;
let mut recorded = rows
@@ -239,7 +220,7 @@ mod tests {
}
#[tokio::test]
async fn overview_aggregates_llm_calls_and_normalizes_provider_usage() {
async fn overview_aggregates_llm_calls_and_normalizes_token_usage() {
let directory = tempfile::tempdir().unwrap();
let store = Store::connect(&format!(
"sqlite://{}",
@@ -263,7 +244,7 @@ mod tests {
)
.await;
let overview = store.overview(None, None, None, None).await.unwrap();
let overview = store.overview(None, None, None).await.unwrap();
assert_eq!(overview.metrics.llm_calls, 3);
assert_eq!(overview.metrics.successful_calls, 2);
assert_eq!(overview.metrics.failed_calls, 1);
@@ -304,7 +285,7 @@ mod tests {
.await;
let overview = store
.overview(Some(now - 1_000), Some(now + 1_000), None, None)
.overview(Some(now - 1_000), Some(now + 1_000), None)
.await
.unwrap();
@@ -325,7 +306,6 @@ mod tests {
Some(now - 1_000),
Some(now + 1_000),
Some(r#"["missing-model"]"#),
None,
)
.await
.unwrap();
File diff suppressed because it is too large Load Diff
+29 -4
View File
@@ -69,12 +69,37 @@ impl Store {
#[cfg(test)]
mod tests {
use super::*;
use crate::model::{ModelConfigInput, ModelType, OPENAI_CHAT_ENDPOINT};
#[tokio::test]
async fn clears_observability_without_removing_configuration() {
let store = Store::connect("sqlite::memory:").await.unwrap();
sqlx::query("INSERT INTO provider_endpoints(name, provider_type, base_url, api_key, created_at_ms, updated_at_ms) VALUES ('Example', 'openai-chat', 'https://example.com', 'secret', 1, 1)")
.execute(store.pool()).await.unwrap();
store
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Model".into(),
model_id: "model".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
sqlx::query("INSERT INTO llm_calls(call_id, run_id, conversation_id, provider_call_index, provider_type, provider_url, request_type, request_url, model_id, display_name, status, created_at_ms, message_count, tool_count, detailed) VALUES ('call-1', 'run-1', 'conversation-1', 0, 'openai-chat', 'https://example.com', 'openai-chat', 'https://example.com/v1/chat/completions', 'model', 'Model', 'completed', 1, 1, 0, 0)")
.execute(store.pool()).await.unwrap();
@@ -82,11 +107,11 @@ mod tests {
let cleared = store.clear_statistics_storage().await.unwrap();
assert_eq!(cleared.bytes, 0);
assert_eq!(cleared.call_count, 0);
let provider_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM provider_endpoints")
let model_count: i64 = sqlx::query_scalar("SELECT COUNT(*) FROM model_configs")
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(provider_count, 1);
assert_eq!(model_count, 1);
store
.record_llm_request(
+25 -29
View File
@@ -9,8 +9,8 @@ use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{connect, proto::agent::v1 as pb, CursorCommand, CursorSessionRegistry},
model::{
ContentPart, ConversationId, MessageContent, Origin, ProjectedContent,
ProviderEndpointInput, ProviderModelInput, ProviderType, Role, Usage,
ContentPart, ConversationId, MessageContent, ModelConfigInput, ModelType, Origin,
ProjectedContent, Role, Usage, OPENAI_CHAT_ENDPOINT,
},
provider::{FinishReason, ModelEvent},
};
@@ -19,34 +19,30 @@ use prost::Message;
#[tokio::test]
async fn summarize_replaces_model_history_and_preserves_cursor_history() {
let (_directory, store) = fixtures::temp_store().await;
let endpoint = store
.create_provider(&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiChat,
base_url: "https://example.com/v1".into(),
api_key: None,
custom_headers: serde_json::json!({}),
extra_params: serde_json::json!({}),
})
.await
.unwrap();
let model = store
.save_provider_model(
endpoint.provider_id,
&ProviderModelInput {
model_id: "test-model".into(),
display_name: "Test Model".into(),
endpoint_type: ProviderType::OpenAiChat,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: false,
},
)
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Test Model".into(),
model_id: "test-model".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();
@@ -1,7 +1,7 @@
use cursor_server::{
model::{
ConversationId, ModelSpec, NewLlmCall, PreparedRun, PromptSpec, ProviderEndpointInput,
ProviderModelInput, ProviderType, RunAction, RunId, RunKind, Usage,
ConversationId, ModelConfigInput, ModelSpec, ModelType, NewLlmCall, PreparedRun,
PromptSpec, ProviderType, RunAction, RunId, RunKind, Usage, OPENAI_CHAT_ENDPOINT,
},
store::{RunStatus, Store},
};
@@ -13,107 +13,91 @@ async fn store() -> (tempfile::TempDir, Store) {
(directory, store)
}
#[tokio::test]
async fn provider_secret_is_write_only_and_model_hash_is_stable() {
let (_directory, store) = store().await;
let provider = store
.create_provider(&ProviderEndpointInput {
name: "Local".into(),
provider_type: ProviderType::OpenAiChat,
base_url: "https://example.com/v1/".into(),
api_key: Some("secret".into()),
custom_headers: serde_json::json!({"x-route":"one", "authorization":"header-secret"}),
extra_params: serde_json::json!({"temperature":0}),
})
.await
.unwrap();
assert!(provider.has_api_key);
assert!(!serde_json::to_string(&provider).unwrap().contains("secret"));
assert_eq!(
provider.custom_headers["authorization"],
serde_json::Value::Null
);
let updated = store
.update_provider(
provider.provider_id,
&ProviderEndpointInput {
name: "Renamed".into(),
provider_type: provider.provider_type,
base_url: provider.base_url.clone(),
api_key: None,
custom_headers: provider.custom_headers.clone(),
extra_params: provider.extra_params.clone(),
},
)
.await
.unwrap();
assert_eq!(updated.name, "Renamed");
assert_eq!(
store
.provider(provider.provider_id)
.await
.unwrap()
.unwrap()
.custom_headers["authorization"],
"header-secret"
);
fn model_input() -> ModelConfigInput {
ModelConfigInput {
sort_order: 0,
display_name: "Model A".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "secret".into(),
tooltip_data: "Model A".into(),
model_id: "model-a".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: true,
openai_extra_params: serde_json::json!({"temperature":0}),
custom_headers_enabled: true,
custom_headers: serde_json::json!({"x-route":"one"}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
}
}
let model = store
.save_provider_model(
provider.provider_id,
&ProviderModelInput {
model_id: "model-a".into(),
display_name: "Model A".into(),
endpoint_type: ProviderType::OpenAiChat,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: true,
},
)
.await
.unwrap();
assert_eq!(model.model_hash, "bab5019a");
assert!(model.supports_image_generation);
#[tokio::test]
async fn model_configuration_round_trips_and_hash_uses_v0049_identity() {
let (_directory, store) = store().await;
let model = store.create_model(&model_input()).await.unwrap();
assert_eq!(model.model_hash.len(), 16);
assert_eq!(model.api_key, "secret");
assert_eq!(model.custom_headers["x-route"], "one");
let original_hash = model.model_hash.clone();
let mut input = model_input();
input.base_url = "https://example.com/v1/chat/completions".into();
input.sort_order = 3;
input.tooltip_data = "Updated tooltip".into();
let updated = store.update_model(&original_hash, &input).await.unwrap();
assert_eq!(updated.model_hash, original_hash);
assert_eq!(updated.sort_order, 3);
assert_eq!(updated.tooltip_data, "Updated tooltip");
}
#[tokio::test]
async fn arbitrary_request_url_is_independent_from_openai_protocol() {
let (_directory, store) = store().await;
let mut input = model_input();
input.base_url = "https://proxy.example.com/arbitrary/generate?api-version=2026-01-01".into();
let chat = store.create_model(&input).await.unwrap();
assert_eq!(chat.provider_type(), ProviderType::OpenAiChat);
assert_eq!(chat.request_url().unwrap(), input.base_url);
input.openai_endpoint = "/v1/responses".into();
let responses = store.update_model(&chat.model_hash, &input).await.unwrap();
assert_eq!(responses.provider_type(), ProviderType::OpenAiResponses);
assert_eq!(responses.request_url().unwrap(), input.base_url);
}
#[tokio::test]
async fn standard_server_address_resolves_to_the_same_model_identity_as_a_complete_url() {
let (_directory, store) = store().await;
let complete = model_input();
let mut standard = complete.clone();
standard.base_url = "https://example.com/v1".into();
standard.use_full_url = false;
let model = store.create_model(&standard).await.unwrap();
assert!(!model.use_full_url);
assert_eq!(
model.request_url().unwrap(),
"https://example.com/v1/chat/completions"
);
assert_eq!(
model.model_hash,
cursor_server::model::model_hash(&complete).unwrap()
);
}
#[tokio::test]
async fn call_summary_is_always_stored_and_payloads_follow_detailed_setting() {
let (_directory, store) = store().await;
let provider = store
.create_provider(&ProviderEndpointInput {
name: "Local".into(),
provider_type: ProviderType::OpenAiChat,
base_url: "https://example.com/v1".into(),
api_key: None,
custom_headers: serde_json::json!({}),
extra_params: serde_json::json!({}),
})
.await
.unwrap();
let model = store
.save_provider_model(
provider.provider_id,
&ProviderModelInput {
model_id: "model-a".into(),
display_name: "Model A".into(),
endpoint_type: ProviderType::OpenAiChat,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: false,
},
)
.await
.unwrap();
let model = store.create_model(&model_input()).await.unwrap();
let call = NewLlmCall {
call_id: "call-1".into(),
run_id: "run-1".into(),
@@ -121,7 +105,7 @@ async fn call_summary_is_always_stored_and_payloads_follow_detailed_setting() {
provider_call_index: 0,
model_hash: model.model_hash,
provider_type: ProviderType::OpenAiChat,
provider_url: provider.base_url,
provider_url: model.base_url.clone(),
request_type: ProviderType::OpenAiChat,
request_url: "https://example.com/v1/chat/completions".into(),
model_id: model.model_id,
+27 -31
View File
@@ -3,8 +3,8 @@ use std::time::Duration;
use axum::{http::header, response::IntoResponse, routing::post, Router};
use cursor_server::{
model::{
ModelInvocation, ModelRequest, ModelSpec, PromptSpec, ProviderEndpointInput,
ProviderModelInput, ProviderType,
ModelConfigInput, ModelInvocation, ModelRequest, ModelSpec, ModelType, PromptSpec,
OPENAI_CHAT_ENDPOINT,
},
provider::{ModelEvent, Provider, ProviderRouter},
store::Store,
@@ -105,7 +105,7 @@ async fn cursor_trace_links_detailed_artifacts_to_the_logical_run() {
#[tokio::test]
async fn records_one_summary_and_raw_payloads_for_one_provider_request() {
let app = Router::new().route(
"/v1/chat/completions",
"/proxy/generate",
post(|| async {
(
[(header::CONTENT_TYPE, "text/event-stream")],
@@ -124,34 +124,30 @@ async fn records_one_summary_and_raw_payloads_for_one_provider_request() {
let (_directory, store) = test_store("observability.db").await;
store.set_detailed_logging(true).await.unwrap();
let endpoint = store
.create_provider(&ProviderEndpointInput {
name: "test".into(),
provider_type: ProviderType::OpenAiChat,
base_url: format!("http://{address}/v1"),
api_key: Some("not-recorded".into()),
custom_headers: serde_json::json!({"x-safe":"visible","authorization":"hidden"}),
extra_params: serde_json::json!({}),
})
.await
.unwrap();
let model = store
.save_provider_model(
endpoint.provider_id,
&ProviderModelInput {
model_id: "actual-model".into(),
display_name: "Display Model".into(),
endpoint_type: ProviderType::OpenAiChat,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: false,
},
)
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Display Model".into(),
model_type: ModelType::OpenAi,
base_url: format!("http://{address}/proxy/generate"),
use_full_url: true,
api_key: "not-recorded".into(),
tooltip_data: "Display Model".into(),
model_id: "actual-model".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: true,
custom_headers: serde_json::json!({"x-safe":"visible","authorization":"hidden"}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let provider = ProviderRouter::new(store.clone(), Duration::from_secs(5));
@@ -183,7 +179,7 @@ async fn records_one_summary_and_raw_payloads_for_one_provider_request() {
let call = store.llm_call("call-1").await.unwrap().unwrap();
assert_eq!(call.status, "completed");
assert_eq!(call.request_type, "openai-chat");
assert!(call.request_url.ends_with("/v1/chat/completions"));
assert_eq!(call.request_url, format!("http://{address}/proxy/generate"));
assert_eq!(call.total_tokens, Some(12));
assert!(call.ttfb_ms.is_some());
assert!(call.ttft_ms.is_some());
+171
View File
@@ -41,3 +41,174 @@ async fn version_two_database_upgrades_with_cursor_request_mapping() {
.iter()
.any(|column| column.get::<String, _>("name") == "cursor_request_id"));
}
#[tokio::test]
async fn provider_and_model_rows_upgrade_to_flat_model_configuration() {
let directory = tempfile::tempdir().unwrap();
let database = directory.path().join("flat-model-upgrade.db");
let pool = sqlx::SqlitePool::connect_with(
SqliteConnectOptions::new()
.filename(&database)
.create_if_missing(true)
.foreign_keys(true),
)
.await
.unwrap();
let all = sqlx::migrate!("./migrations");
let prior = Migrator {
migrations: Cow::Owned(
all.iter()
.filter(|migration| migration.version <= 3)
.cloned()
.collect(),
),
ignore_missing: false,
locking: true,
no_tx: false,
};
prior.run(&pool).await.unwrap();
sqlx::query(
r#"INSERT INTO provider_endpoints(
name, provider_type, base_url, api_key, custom_headers_json, extra_params_json,
created_at_ms, updated_at_ms
) VALUES ('Example', 'openai-chat', 'https://example.com/v1', 'secret',
'{"x-client":"cursor-byok"}', '{"service_tier":"priority"}', 10, 11)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"INSERT INTO provider_models(
model_hash, provider_id, model_id, display_name, endpoint_type, request_url,
enabled, sort_order, reasoning_enabled, supports_image_generation,
created_at_ms, updated_at_ms
) VALUES
('rspns001', 1, 'model-b', 'Model B', 'openai-responses',
'https://proxy.example.com/arbitrary/generate?api-version=2026-01-01',
1, 5, 0, 0, 12, 13),
('anthr001', 1, 'model-c', 'Model C', 'anthropic', '/proxy/claude',
1, 6, 0, 0, 12, 13),
('anthstd1', 1, 'model-d', 'Model D', 'anthropic', '',
1, 7, 0, 0, 12, 13)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"INSERT INTO provider_models(
model_hash, provider_id, model_id, display_name, endpoint_type, request_url,
enabled, sort_order, context_window_tokens, max_output_tokens,
reasoning_enabled, reasoning_effort, supports_image_generation,
created_at_ms, updated_at_ms
) VALUES ('12345678', 1, 'model-a', 'Model A', 'openai-chat', '',
1, 4, 200000, 8192, 1, 'high', 0, 12, 13)"#,
)
.execute(&pool)
.await
.unwrap();
sqlx::query(
r#"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 ('call-1', 'run-1', 'conversation-1', 0, '12345678',
'openai-chat', 'https://example.com/v1', 'openai-chat',
'https://example.com/v1/chat/completions', 'model-a', 'Model A',
'completed', 14, 1, 0, 0)"#,
)
.execute(&pool)
.await
.unwrap();
all.run(&pool).await.unwrap();
drop(pool);
let store = Store::connect(&format!("sqlite://{}", database.display()))
.await
.unwrap();
let row = sqlx::query(
r#"SELECT model_hash, sort_order, display_name, model_type, base_url, api_key,
tooltip_data, model_id, reasoning_effort, openai_endpoint, use_full_url,
openai_extra_params_enabled, openai_extra_params_json,
custom_headers_enabled, custom_headers_json, context_window_tokens,
max_completion_tokens
FROM model_configs WHERE model_hash = '12345678'"#,
)
.fetch_one(store.pool())
.await
.unwrap();
assert_eq!(row.get::<String, _>("model_hash"), "12345678");
assert_eq!(row.get::<i64, _>("sort_order"), 4);
assert_eq!(row.get::<String, _>("model_type"), "openai");
assert_eq!(row.get::<String, _>("base_url"), "https://example.com/v1");
assert_eq!(row.get::<String, _>("api_key"), "secret");
assert_eq!(row.get::<i64, _>("use_full_url"), 0);
assert_eq!(row.get::<String, _>("reasoning_effort"), "high");
assert_eq!(
row.get::<String, _>("openai_endpoint"),
"/v1/chat/completions"
);
assert_eq!(row.get::<i64, _>("openai_extra_params_enabled"), 1);
assert_eq!(
row.get::<String, _>("openai_extra_params_json"),
r#"{"service_tier":"priority"}"#
);
assert_eq!(row.get::<i64, _>("custom_headers_enabled"), 1);
assert_eq!(row.get::<i64, _>("context_window_tokens"), 200000);
assert_eq!(row.get::<i64, _>("max_completion_tokens"), 8192);
let migrated_rows = sqlx::query(
"SELECT model_hash, model_type, base_url, use_full_url, openai_endpoint FROM model_configs WHERE model_hash IN ('rspns001', 'anthr001', 'anthstd1') ORDER BY model_hash",
)
.fetch_all(store.pool())
.await
.unwrap();
assert_eq!(migrated_rows[0].get::<String, _>("model_hash"), "anthr001");
assert_eq!(migrated_rows[0].get::<String, _>("model_type"), "anthropic");
assert_eq!(
migrated_rows[0].get::<String, _>("base_url"),
"https://example.com/v1/proxy/claude"
);
assert_eq!(migrated_rows[0].get::<i64, _>("use_full_url"), 1);
assert_eq!(migrated_rows[0].get::<String, _>("openai_endpoint"), "");
assert_eq!(migrated_rows[1].get::<String, _>("model_hash"), "anthstd1");
assert_eq!(migrated_rows[1].get::<String, _>("model_type"), "anthropic");
assert_eq!(
migrated_rows[1].get::<String, _>("base_url"),
"https://example.com/v1"
);
assert_eq!(migrated_rows[1].get::<i64, _>("use_full_url"), 0);
assert_eq!(migrated_rows[1].get::<String, _>("openai_endpoint"), "");
assert_eq!(migrated_rows[2].get::<String, _>("model_hash"), "rspns001");
assert_eq!(migrated_rows[2].get::<String, _>("model_type"), "openai");
assert_eq!(
migrated_rows[2].get::<String, _>("base_url"),
"https://proxy.example.com/arbitrary/generate?api-version=2026-01-01"
);
assert_eq!(migrated_rows[2].get::<i64, _>("use_full_url"), 1);
assert_eq!(
migrated_rows[2].get::<String, _>("openai_endpoint"),
"/v1/responses"
);
assert_eq!(
sqlx::query_scalar::<_, String>(
"SELECT model_hash FROM llm_calls WHERE call_id = 'call-1'"
)
.fetch_one(store.pool())
.await
.unwrap(),
"12345678"
);
assert!(sqlx::query("SELECT 1 FROM provider_endpoints")
.fetch_one(store.pool())
.await
.is_err());
assert!(sqlx::query("SELECT 1 FROM provider_models")
.fetch_one(store.pool())
.await
.is_err());
assert!(sqlx::query("PRAGMA foreign_key_check")
.fetch_all(store.pool())
.await
.unwrap()
.is_empty());
}
+24 -30
View File
@@ -12,9 +12,7 @@ use cursor_server::{
cursor::prompting::{PromptAssets, PromptCompiler},
cursor::{connect, proto::agent::v1 as pb},
cursor::{CursorCommand, CursorSessionRegistry},
model::{
ProjectedContent, ProviderEndpointInput, ProviderModelInput, ProviderType, Role, Usage,
},
model::{ModelConfigInput, ModelType, ProjectedContent, Role, Usage, OPENAI_CHAT_ENDPOINT},
provider::{FinishReason, ModelEvent},
};
use prost::Message;
@@ -22,34 +20,30 @@ use prost::Message;
#[tokio::test]
async fn text_turn_runs_from_bidi_request_through_checkpoint_and_end_stream() {
let (_directory, store) = fixtures::temp_store().await;
let endpoint = store
.create_provider(&ProviderEndpointInput {
name: "Test".into(),
provider_type: ProviderType::OpenAiChat,
base_url: "https://example.com/v1".into(),
api_key: None,
custom_headers: serde_json::json!({}),
extra_params: serde_json::json!({}),
})
.await
.unwrap();
let configured_model = store
.save_provider_model(
endpoint.provider_id,
&ProviderModelInput {
model_id: "test-model".into(),
display_name: "Test Model".into(),
endpoint_type: ProviderType::OpenAiChat,
request_url: String::new(),
enabled: true,
sort_order: 0,
context_window_tokens: None,
max_output_tokens: None,
reasoning_enabled: false,
reasoning_effort: None,
supports_image_generation: false,
},
)
.create_model(&ModelConfigInput {
sort_order: 0,
display_name: "Test Model".into(),
model_type: ModelType::OpenAi,
base_url: "https://example.com/v1/chat/completions".into(),
use_full_url: true,
api_key: "test-key".into(),
tooltip_data: "Test Model".into(),
model_id: "test-model".into(),
reasoning_effort: None,
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
openai_extra_params_enabled: false,
openai_extra_params: serde_json::json!({}),
custom_headers_enabled: false,
custom_headers: serde_json::json!({}),
anthropic_extra_params_enabled: false,
anthropic_extra_params: serde_json::json!({}),
context_window_tokens: None,
max_completion_tokens: None,
anthropic_max_tokens: None,
anthropic_thinking_effort: None,
thinking_budget_tokens: None,
})
.await
.unwrap();
let provider = fake_provider::FakeProvider::default();