mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-05 03:56:45 +08:00
merge old config
This commit is contained in:
@@ -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);
|
||||
@@ -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,
|
||||
|
||||
@@ -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
@@ -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(),
|
||||
)
|
||||
|
||||
@@ -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?))
|
||||
}
|
||||
|
||||
@@ -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?,
|
||||
))
|
||||
}
|
||||
|
||||
@@ -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
@@ -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();
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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)?;
|
||||
|
||||
@@ -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?;
|
||||
|
||||
@@ -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"
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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::*;
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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)
|
||||
|
||||
@@ -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);
|
||||
}
|
||||
}
|
||||
@@ -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,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;
|
||||
|
||||
@@ -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(¤t.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
|
||||
);
|
||||
}
|
||||
}
|
||||
@@ -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
@@ -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
@@ -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,
|
||||
@@ -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());
|
||||
|
||||
@@ -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
@@ -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();
|
||||
|
||||
Reference in New Issue
Block a user