mirror of
https://wget.la/https://github.com/leookun/cursor-byok
synced 2026-10-04 02:52:55 +08:00
904 lines
31 KiB
Rust
904 lines
31 KiB
Rust
#[path = "support/fake_provider.rs"]
|
|
mod fake_provider;
|
|
#[path = "support/fixtures.rs"]
|
|
mod fixtures;
|
|
|
|
use std::{collections::HashSet, sync::Arc, time::Duration};
|
|
|
|
use axum::{
|
|
body::{to_bytes, Body},
|
|
http::{header, HeaderMap, Request, StatusCode},
|
|
response::IntoResponse,
|
|
routing::post,
|
|
Json, Router,
|
|
};
|
|
use cursor_server::{
|
|
api::byok,
|
|
control,
|
|
model::{
|
|
ModelConfigInput, ModelType, ProjectedContent, Usage, OPENAI_CHAT_ENDPOINT,
|
|
OPENAI_RESPONSES_ENDPOINT,
|
|
},
|
|
network::NetworkClients,
|
|
plugin::{PluginRegistry, PluginRuntime},
|
|
provider::{FinishReason, ModelEvent, ProviderRouter},
|
|
store::ExternalApiSettings,
|
|
};
|
|
use serde_json::{json, Value};
|
|
use tower::ServiceExt;
|
|
|
|
fn model_input() -> ModelConfigInput {
|
|
ModelConfigInput {
|
|
sort_order: 0,
|
|
display_name: "Example".into(),
|
|
group_name: Some("work".into()),
|
|
model_type: ModelType::OpenAi,
|
|
base_url: "https://example.com/v1".into(),
|
|
use_full_url: false,
|
|
api_key: "upstream".into(),
|
|
tooltip_data: "test".into(),
|
|
model_id: "qwen/model".into(),
|
|
reasoning_effort: None,
|
|
openai_endpoint: OPENAI_CHAT_ENDPOINT.into(),
|
|
openai_extra_params_enabled: false,
|
|
openai_extra_params: json!({}),
|
|
custom_headers_enabled: false,
|
|
custom_headers: json!({}),
|
|
anthropic_extra_params_enabled: false,
|
|
anthropic_extra_params: json!({}),
|
|
context_window_tokens: None,
|
|
max_completion_tokens: None,
|
|
anthropic_max_tokens: None,
|
|
anthropic_thinking_effort: None,
|
|
thinking_budget_tokens: None,
|
|
}
|
|
}
|
|
|
|
async fn setup() -> (
|
|
axum::Router,
|
|
fake_provider::FakeProvider,
|
|
cursor_server::store::Store,
|
|
tempfile::TempDir,
|
|
) {
|
|
let (directory, store) = fixtures::temp_store().await;
|
|
store.create_model(&model_input()).await.unwrap();
|
|
let runtime = PluginRuntime::managed().unwrap();
|
|
let plugins = PluginRegistry::managed(store.clone(), runtime.clone(), "0.1.0".into()).unwrap();
|
|
let provider = fake_provider::FakeProvider::default();
|
|
let shared_provider = Arc::new(provider.clone());
|
|
let control = control::ControlService::new(
|
|
store.clone(),
|
|
shared_provider.clone(),
|
|
runtime,
|
|
plugins.clone(),
|
|
NetworkClients::new(store.clone()),
|
|
"0.1.0".into(),
|
|
)
|
|
.unwrap();
|
|
let router = byok::router(store.clone(), plugins, shared_provider, None)
|
|
.merge(control::api_router(control));
|
|
(router, provider, store, directory)
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn management_settings_enable_the_external_route_without_restart() {
|
|
let (router, _provider, _store, _directory) = setup().await;
|
|
let (status, body) = send(
|
|
router.clone(),
|
|
"GET",
|
|
"/__byok-api__/api/settings/external-api",
|
|
None,
|
|
json!({}),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(
|
|
serde_json::from_str::<Value>(&body).unwrap()["enabled"],
|
|
false
|
|
);
|
|
let (status, _) = send(
|
|
router.clone(),
|
|
"PUT",
|
|
"/__byok-api__/api/settings/external-api",
|
|
None,
|
|
json!({"enabled":true,"api_key":"changed-key"}),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(
|
|
send(
|
|
router,
|
|
"GET",
|
|
"/byok/v1/models",
|
|
Some("changed-key"),
|
|
json!({})
|
|
)
|
|
.await
|
|
.0,
|
|
StatusCode::OK
|
|
);
|
|
}
|
|
|
|
async fn send(
|
|
router: axum::Router,
|
|
method: &str,
|
|
path: &str,
|
|
key: Option<&str>,
|
|
body: Value,
|
|
) -> (StatusCode, String) {
|
|
let mut request = Request::builder()
|
|
.method(method)
|
|
.uri(path)
|
|
.header(header::CONTENT_TYPE, "application/json");
|
|
if let Some(key) = key {
|
|
request = request.header(header::AUTHORIZATION, format!("Bearer {key}"));
|
|
}
|
|
let response = router
|
|
.oneshot(request.body(Body::from(body.to_string())).unwrap())
|
|
.await
|
|
.unwrap();
|
|
let status = response.status();
|
|
let bytes = to_bytes(response.into_body(), 1024 * 1024).await.unwrap();
|
|
(status, String::from_utf8(bytes.to_vec()).unwrap())
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn disabled_and_unauthorized_requests_cannot_list_models() {
|
|
let (router, _provider, store, _directory) = setup().await;
|
|
assert_eq!(
|
|
send(
|
|
router.clone(),
|
|
"GET",
|
|
"/byok/v1/models",
|
|
Some("secret"),
|
|
json!({})
|
|
)
|
|
.await
|
|
.0,
|
|
StatusCode::FORBIDDEN
|
|
);
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
assert_eq!(
|
|
send(router.clone(), "GET", "/byok/v1/models", None, json!({}))
|
|
.await
|
|
.0,
|
|
StatusCode::UNAUTHORIZED
|
|
);
|
|
assert_eq!(
|
|
send(
|
|
router.clone(),
|
|
"GET",
|
|
"/byok/v1/models",
|
|
Some("wrong"),
|
|
json!({})
|
|
)
|
|
.await
|
|
.0,
|
|
StatusCode::UNAUTHORIZED
|
|
);
|
|
let (status, body) = send(router, "GET", "/byok/v1/models", Some("secret"), json!({})).await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert_eq!(
|
|
serde_json::from_str::<Value>(&body).unwrap()["data"][0]["id"],
|
|
"work/qwen/model"
|
|
);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn duplicate_public_model_ids_use_the_first_configured_model() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
let first = store.models().await.unwrap().remove(0);
|
|
let mut duplicate = model_input();
|
|
duplicate.sort_order = 1;
|
|
duplicate.display_name = "Backup".into();
|
|
duplicate.base_url = "https://backup.example.com/v1".into();
|
|
store.create_model(&duplicate).await.unwrap();
|
|
let mut distinct = model_input();
|
|
distinct.sort_order = 2;
|
|
distinct.model_id = "qwen/other".into();
|
|
store.create_model(&distinct).await.unwrap();
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
|
|
let (status, body) = send(
|
|
router.clone(),
|
|
"GET",
|
|
"/byok/v1/models",
|
|
Some("secret"),
|
|
json!({}),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK, "{body}");
|
|
let models = serde_json::from_str::<Value>(&body).unwrap();
|
|
let ids = models["data"]
|
|
.as_array()
|
|
.unwrap()
|
|
.iter()
|
|
.map(|model| model["id"].as_str().unwrap())
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(ids, ["work/qwen/model", "work/qwen/other"]);
|
|
|
|
provider.push(vec![
|
|
ModelEvent::TextDelta("ok".into()),
|
|
ModelEvent::Done(FinishReason::Stop),
|
|
]);
|
|
let (status, body) = send(
|
|
router,
|
|
"POST",
|
|
"/byok/v1/chat/completions",
|
|
Some("secret"),
|
|
json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK, "{body}");
|
|
assert_eq!(provider.requests()[0].model.model_id, first.model_hash);
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn all_three_protocols_use_the_public_model_id() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
for (path, body) in [
|
|
(
|
|
"/byok/v1/chat/completions",
|
|
json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}),
|
|
),
|
|
(
|
|
"/byok/v1/responses",
|
|
json!({"model":"work/qwen/model","input":"hello"}),
|
|
),
|
|
(
|
|
"/byok/v1/messages",
|
|
json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}]}),
|
|
),
|
|
] {
|
|
provider.push(vec![
|
|
ModelEvent::TextStart,
|
|
ModelEvent::TextDelta("world".into()),
|
|
ModelEvent::TextEnd,
|
|
ModelEvent::Done(FinishReason::Stop),
|
|
]);
|
|
let (status, response) = send(router.clone(), "POST", path, Some("secret"), body).await;
|
|
assert_eq!(status, StatusCode::OK, "{path}: {response}");
|
|
assert!(response.contains("world"), "{path}: {response}");
|
|
}
|
|
assert_eq!(provider.requests().len(), 3);
|
|
assert!(provider
|
|
.requests()
|
|
.iter()
|
|
.all(|request| request.model.model_id != "work/qwen/model"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn streaming_chat_returns_incremental_sse_and_tool_calls() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
provider.push(vec![
|
|
ModelEvent::TextStart,
|
|
ModelEvent::TextDelta("hello".into()),
|
|
ModelEvent::TextEnd,
|
|
ModelEvent::ToolCallStart {
|
|
index: 0,
|
|
call_id: "call_1".into(),
|
|
name: "lookup".into(),
|
|
},
|
|
ModelEvent::ToolCallArgumentsDelta {
|
|
index: 0,
|
|
delta: "{\"q\":1}".into(),
|
|
},
|
|
ModelEvent::ToolCallEnd { index: 0 },
|
|
ModelEvent::Done(FinishReason::ToolUse),
|
|
]);
|
|
let (status, body) = send(router, "POST", "/byok/v1/chat/completions", Some("secret"),
|
|
json!({"model":"work/qwen/model","stream":true,"messages":[{"role":"user","content":"hello"}]})).await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert!(body.contains("chat.completion.chunk"));
|
|
assert!(body.contains("lookup"));
|
|
assert!(body.contains("tool_calls"));
|
|
assert!(body.contains("[DONE]"));
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn chat_tool_result_keeps_the_assistant_function_name() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
provider.push(vec![
|
|
ModelEvent::TextDelta("done".into()),
|
|
ModelEvent::Done(FinishReason::Stop),
|
|
]);
|
|
let body = json!({"model":"work/qwen/model","messages":[
|
|
{"role":"user","content":"find it"},
|
|
{"role":"assistant","tool_calls":[{"id":"call_1","type":"function","function":{"name":"lookup","arguments":"{\"q\":1}"}}]},
|
|
{"role":"tool","tool_call_id":"call_1","content":"found"}
|
|
]});
|
|
assert_eq!(
|
|
send(
|
|
router,
|
|
"POST",
|
|
"/byok/v1/chat/completions",
|
|
Some("secret"),
|
|
body
|
|
)
|
|
.await
|
|
.0,
|
|
StatusCode::OK
|
|
);
|
|
let requests = provider.requests();
|
|
let ProjectedContent::ToolResult(result) = &requests[0].history[2].content else {
|
|
panic!("expected tool result");
|
|
};
|
|
assert_eq!(result.name, "lookup");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn responses_and_messages_stream_with_protocol_end_events() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
for (path, request, terminal) in [
|
|
(
|
|
"/byok/v1/responses",
|
|
json!({"model":"work/qwen/model","input":"hello","stream":true}),
|
|
"response.completed",
|
|
),
|
|
(
|
|
"/byok/v1/messages",
|
|
json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"stream":true}),
|
|
"message_stop",
|
|
),
|
|
] {
|
|
provider.push(vec![
|
|
ModelEvent::TextDelta("world".into()),
|
|
ModelEvent::Done(FinishReason::Stop),
|
|
]);
|
|
let (status, body) = send(router.clone(), "POST", path, Some("secret"), request).await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
assert!(body.contains(terminal), "{path}: {body}");
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn responses_stream_emits_complete_text_and_tool_item_lifecycles() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
provider.push(vec![
|
|
ModelEvent::TextStart,
|
|
ModelEvent::TextDelta("hello".into()),
|
|
ModelEvent::TextEnd,
|
|
ModelEvent::ToolCallStart {
|
|
index: 0,
|
|
call_id: "call_1".into(),
|
|
name: "lookup".into(),
|
|
},
|
|
ModelEvent::ToolCallArgumentsDelta {
|
|
index: 0,
|
|
delta: "{\"q\":1}".into(),
|
|
},
|
|
ModelEvent::ToolCallEnd { index: 0 },
|
|
ModelEvent::Done(FinishReason::ToolUse),
|
|
]);
|
|
let (status, body) = send(
|
|
router,
|
|
"POST",
|
|
"/byok/v1/responses",
|
|
Some("secret"),
|
|
json!({"model":"work/qwen/model","input":"hello","stream":true}),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
let events = body
|
|
.lines()
|
|
.filter_map(|line| line.strip_prefix("event: "))
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(
|
|
events,
|
|
[
|
|
"response.created",
|
|
"response.in_progress",
|
|
"response.output_item.added",
|
|
"response.content_part.added",
|
|
"response.output_text.delta",
|
|
"response.output_text.done",
|
|
"response.content_part.done",
|
|
"response.output_item.done",
|
|
"response.output_item.added",
|
|
"response.function_call_arguments.delta",
|
|
"response.function_call_arguments.done",
|
|
"response.output_item.done",
|
|
"response.completed",
|
|
]
|
|
);
|
|
assert!(body.contains("\"output_index\":1"), "{body}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn messages_stream_closes_each_content_block_before_stopping() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
provider.push(vec![
|
|
ModelEvent::TextStart,
|
|
ModelEvent::TextDelta("hello".into()),
|
|
ModelEvent::TextEnd,
|
|
ModelEvent::ToolCallStart {
|
|
index: 0,
|
|
call_id: "call_1".into(),
|
|
name: "lookup".into(),
|
|
},
|
|
ModelEvent::ToolCallArgumentsDelta {
|
|
index: 0,
|
|
delta: "{\"q\":1}".into(),
|
|
},
|
|
ModelEvent::ToolCallEnd { index: 0 },
|
|
ModelEvent::Done(FinishReason::ToolUse),
|
|
]);
|
|
let (status, body) = send(router, "POST", "/byok/v1/messages", Some("secret"),
|
|
json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}],"stream":true})).await;
|
|
assert_eq!(status, StatusCode::OK);
|
|
let events = body
|
|
.lines()
|
|
.filter_map(|line| line.strip_prefix("event: "))
|
|
.collect::<Vec<_>>();
|
|
assert_eq!(
|
|
events,
|
|
[
|
|
"message_start",
|
|
"content_block_start",
|
|
"content_block_delta",
|
|
"content_block_stop",
|
|
"content_block_start",
|
|
"content_block_delta",
|
|
"content_block_stop",
|
|
"message_delta",
|
|
"message_stop",
|
|
]
|
|
);
|
|
assert!(body.contains("\"index\":1"), "{body}");
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn all_protocols_preserve_cached_usage_in_streaming_and_complete_responses() {
|
|
let (router, provider, store, _directory) = setup().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
let usage = Usage {
|
|
input_tokens: Some(1_000),
|
|
context_input_tokens: Some(1_000),
|
|
output_tokens: Some(50),
|
|
total_tokens: Some(1_050),
|
|
cache_read_tokens: Some(800),
|
|
cache_write_tokens: Some(20),
|
|
reasoning_tokens: Some(10),
|
|
};
|
|
for (path, request, usage_pointer, cached_pointer, expected_input) in [
|
|
(
|
|
"/byok/v1/chat/completions",
|
|
json!({"model":"work/qwen/model","messages":[{"role":"user","content":"hello"}]}),
|
|
"/usage",
|
|
"/prompt_tokens_details/cached_tokens",
|
|
1_000,
|
|
),
|
|
(
|
|
"/byok/v1/responses",
|
|
json!({"model":"work/qwen/model","input":"hello"}),
|
|
"/response/usage",
|
|
"/input_tokens_details/cached_tokens",
|
|
1_000,
|
|
),
|
|
(
|
|
"/byok/v1/messages",
|
|
json!({"model":"work/qwen/model","max_tokens":100,"messages":[{"role":"user","content":"hello"}]}),
|
|
"/usage",
|
|
"/cache_read_input_tokens",
|
|
180,
|
|
),
|
|
] {
|
|
for stream in [false, true] {
|
|
provider.push(vec![
|
|
ModelEvent::TextDelta("world".into()),
|
|
ModelEvent::Usage(usage),
|
|
ModelEvent::Done(FinishReason::Stop),
|
|
]);
|
|
let mut request = request.clone();
|
|
request["stream"] = json!(stream);
|
|
let (status, body) = send(router.clone(), "POST", path, Some("secret"), request).await;
|
|
assert_eq!(status, StatusCode::OK, "{path}: {body}");
|
|
let response = if stream {
|
|
body.lines()
|
|
.filter_map(|line| line.strip_prefix("data: "))
|
|
.filter_map(|line| serde_json::from_str::<Value>(line).ok())
|
|
.find(|event| match path {
|
|
"/byok/v1/chat/completions" => event.get("usage").is_some(),
|
|
"/byok/v1/responses" => event["type"] == "response.completed",
|
|
_ => event["type"] == "message_delta",
|
|
})
|
|
.unwrap_or_else(|| panic!("missing usage event in {path}: {body}"))
|
|
} else {
|
|
serde_json::from_str::<Value>(&body).unwrap()
|
|
};
|
|
let usage = response
|
|
.pointer(if stream { usage_pointer } else { "/usage" })
|
|
.unwrap();
|
|
assert_eq!(
|
|
usage.pointer(cached_pointer),
|
|
Some(&json!(800)),
|
|
"{path} stream={stream}: {body}"
|
|
);
|
|
let input_field = if path == "/byok/v1/chat/completions" {
|
|
"prompt_tokens"
|
|
} else {
|
|
"input_tokens"
|
|
};
|
|
assert_eq!(
|
|
usage[input_field], expected_input,
|
|
"{path} stream={stream}: {body}"
|
|
);
|
|
if path == "/byok/v1/messages" {
|
|
assert_eq!(usage["cache_creation_input_tokens"], 20, "{body}");
|
|
}
|
|
if stream && path == "/byok/v1/chat/completions" {
|
|
let finished = body.find("\"finish_reason\":\"stop\"").unwrap();
|
|
let usage_position = body.find("\"cached_tokens\":800").unwrap();
|
|
let done = body.find("[DONE]").unwrap();
|
|
assert!(finished < usage_position && usage_position < done, "{body}");
|
|
}
|
|
}
|
|
}
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn all_entry_and_upstream_protocol_pairs_preserve_cache_usage() {
|
|
let (_directory, store) = fixtures::temp_store().await;
|
|
store.create_model(&model_input()).await.unwrap();
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
let runtime = PluginRuntime::managed().unwrap();
|
|
let plugins = PluginRegistry::managed(store.clone(), runtime, "0.1.0".into()).unwrap();
|
|
let fake = fake_provider::FakeProvider::default();
|
|
let inner = byok::router(store.clone(), plugins.clone(), Arc::new(fake.clone()), None);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let port = listener.local_addr().unwrap().port();
|
|
let server = tokio::spawn(async move { axum::serve(listener, inner).await.unwrap() });
|
|
|
|
for (order, group, model_type, endpoint) in [
|
|
(1, "bridge-chat", ModelType::OpenAi, OPENAI_CHAT_ENDPOINT),
|
|
(
|
|
2,
|
|
"bridge-responses",
|
|
ModelType::OpenAi,
|
|
OPENAI_RESPONSES_ENDPOINT,
|
|
),
|
|
(3, "bridge-messages", ModelType::Anthropic, ""),
|
|
] {
|
|
let mut model = model_input();
|
|
model.sort_order = order;
|
|
model.display_name = format!("Bridge {group}");
|
|
model.group_name = Some(group.into());
|
|
model.model_type = model_type;
|
|
model.openai_endpoint = endpoint.into();
|
|
model.base_url = format!("http://127.0.0.1:{port}/byok/v1");
|
|
model.model_id = "work/qwen/model".into();
|
|
model.api_key = "secret".into();
|
|
store.create_model(&model).await.unwrap();
|
|
}
|
|
let provider = ProviderRouter::new(
|
|
store.clone(),
|
|
plugins.clone(),
|
|
NetworkClients::new(store.clone()),
|
|
Duration::from_secs(10),
|
|
Duration::from_secs(10),
|
|
);
|
|
let outer = byok::router(
|
|
store.clone(),
|
|
plugins,
|
|
Arc::new(provider),
|
|
Some(byok::NativeForwarder::new(
|
|
store.clone(),
|
|
NetworkClients::new(store.clone()),
|
|
Duration::from_secs(10),
|
|
Duration::from_secs(10),
|
|
)),
|
|
);
|
|
for (path, request) in [
|
|
(
|
|
"/byok/v1/chat/completions",
|
|
json!({"messages":[{"role":"user","content":"hello"}],
|
|
"tools":[{"type":"function","function":{"name":"lookup","description":"Look up a value","parameters":{"type":"object","properties":{"q":{"type":"string"}}}}}]}),
|
|
),
|
|
(
|
|
"/byok/v1/responses",
|
|
json!({"input":"hello",
|
|
"tools":[{"type":"function","name":"lookup","description":"Look up a value","parameters":{"type":"object","properties":{"q":{"type":"string"}}}}]}),
|
|
),
|
|
(
|
|
"/byok/v1/messages",
|
|
json!({"max_tokens":100,"messages":[{"role":"user","content":"hello"}],
|
|
"tools":[{"name":"lookup","description":"Look up a value","input_schema":{"type":"object","properties":{"q":{"type":"string"}}}}]}),
|
|
),
|
|
] {
|
|
for group in ["bridge-chat", "bridge-responses", "bridge-messages"] {
|
|
for stream in [false, true] {
|
|
let prior_calls: HashSet<_> = store
|
|
.llm_calls(100)
|
|
.await
|
|
.unwrap()
|
|
.into_iter()
|
|
.map(|call| call.call_id)
|
|
.collect();
|
|
fake.push(vec![
|
|
ModelEvent::TextStart,
|
|
ModelEvent::TextDelta("world".into()),
|
|
ModelEvent::TextEnd,
|
|
ModelEvent::ToolCallStart {
|
|
index: 0,
|
|
call_id: "call_1".into(),
|
|
name: "lookup".into(),
|
|
},
|
|
ModelEvent::ToolCallArgumentsDelta {
|
|
index: 0,
|
|
delta: "{\"q\":\"x\"}".into(),
|
|
},
|
|
ModelEvent::ToolCallEnd { index: 0 },
|
|
ModelEvent::Usage(Usage {
|
|
input_tokens: Some(1_000),
|
|
context_input_tokens: Some(1_000),
|
|
output_tokens: Some(50),
|
|
total_tokens: Some(1_050),
|
|
cache_read_tokens: Some(800),
|
|
..Usage::default()
|
|
}),
|
|
ModelEvent::Done(FinishReason::ToolUse),
|
|
]);
|
|
let mut request = request.clone();
|
|
request["model"] = json!(format!("{group}/work/qwen/model"));
|
|
request["stream"] = json!(stream);
|
|
let (status, body) =
|
|
send(outer.clone(), "POST", path, Some("secret"), request).await;
|
|
assert_eq!(
|
|
status,
|
|
StatusCode::OK,
|
|
"{path} -> {group} stream={stream}: {body}"
|
|
);
|
|
assert!(
|
|
body.contains("world"),
|
|
"{path} -> {group} stream={stream}: {body}"
|
|
);
|
|
assert!(
|
|
body.contains("lookup") && body.contains("call_1"),
|
|
"{path} -> {group} stream={stream}: {body}"
|
|
);
|
|
let requests = fake.requests();
|
|
let forwarded = requests.last().unwrap();
|
|
assert_eq!(
|
|
forwarded.prompt.tools.len(),
|
|
1,
|
|
"{path} -> {group} stream={stream}"
|
|
);
|
|
assert_eq!(
|
|
forwarded.prompt.tools[0].name, "lookup",
|
|
"{path} -> {group} stream={stream}"
|
|
);
|
|
let calls = store.llm_calls(100).await.unwrap();
|
|
let call = calls
|
|
.iter()
|
|
.find(|call| {
|
|
call.display_name == format!("Bridge {group}")
|
|
&& !prior_calls.contains(&call.call_id)
|
|
})
|
|
.unwrap();
|
|
assert_eq!(
|
|
call.cache_read_tokens,
|
|
Some(800),
|
|
"{path} -> {group} stream={stream}"
|
|
);
|
|
}
|
|
}
|
|
}
|
|
assert_eq!(fake.requests().len(), 18);
|
|
server.abort();
|
|
}
|
|
|
|
#[tokio::test]
|
|
async fn matching_http_protocols_forward_native_requests_and_responses() {
|
|
let (directory, store) = fixtures::temp_store().await;
|
|
store
|
|
.set_external_api_settings(ExternalApiSettings {
|
|
enabled: true,
|
|
api_key: "secret".into(),
|
|
})
|
|
.await
|
|
.unwrap();
|
|
let runtime = PluginRuntime::managed().unwrap();
|
|
let plugins = PluginRegistry::managed(store.clone(), runtime, "0.1.0".into()).unwrap();
|
|
let (sender, mut receiver) = tokio::sync::mpsc::unbounded_channel::<(HeaderMap, Value)>();
|
|
let upstream = Router::new().route(
|
|
"/{*path}",
|
|
post(move |headers: HeaderMap, Json(body): Json<Value>| {
|
|
let sender = sender.clone();
|
|
async move {
|
|
sender.send((headers, body.clone())).unwrap();
|
|
if body["native_error"] == true {
|
|
return (
|
|
StatusCode::TOO_MANY_REQUESTS,
|
|
Json(json!({"error":"native-rate-limit"})),
|
|
)
|
|
.into_response();
|
|
}
|
|
if body["stream"] == true {
|
|
(
|
|
[(header::CONTENT_TYPE, "text/event-stream")],
|
|
"data: {\"native_marker\":\"untouched-stream\"}\n\n",
|
|
)
|
|
.into_response()
|
|
} else {
|
|
Json(json!({"native_marker":"untouched-complete","echo":body})).into_response()
|
|
}
|
|
}
|
|
}),
|
|
);
|
|
let listener = tokio::net::TcpListener::bind("127.0.0.1:0").await.unwrap();
|
|
let port = listener.local_addr().unwrap().port();
|
|
let server = tokio::spawn(async move { axum::serve(listener, upstream).await.unwrap() });
|
|
for (order, group, model_type, endpoint, path, request) in [
|
|
(
|
|
1,
|
|
"chat",
|
|
ModelType::OpenAi,
|
|
OPENAI_CHAT_ENDPOINT,
|
|
"/byok/v1/chat/completions",
|
|
json!({"messages":[{"role":"user","content":"hi"}],"native_extension":{"keep":1}}),
|
|
),
|
|
(
|
|
2,
|
|
"responses",
|
|
ModelType::OpenAi,
|
|
OPENAI_RESPONSES_ENDPOINT,
|
|
"/byok/v1/responses",
|
|
json!({"input":[{"type":"native_unsupported","value":1}],"native_extension":{"keep":2}}),
|
|
),
|
|
(
|
|
3,
|
|
"messages",
|
|
ModelType::Anthropic,
|
|
"",
|
|
"/byok/v1/messages",
|
|
json!({"max_tokens":100,"messages":[{"role":"user","content":[{"type":"document","source":{"type":"url","url":"https://example.com"}}]}],"native_extension":{"keep":3}}),
|
|
),
|
|
] {
|
|
let mut model = model_input();
|
|
model.sort_order = order;
|
|
model.group_name = Some(group.into());
|
|
model.model_type = model_type;
|
|
model.openai_endpoint = endpoint.into();
|
|
model.base_url = format!("http://127.0.0.1:{port}/byok/v1");
|
|
model.model_id = "native-model".into();
|
|
store.create_model(&model).await.unwrap();
|
|
for stream in [false, true] {
|
|
let mut request = request.clone();
|
|
request["model"] = json!(format!("{group}/native-model"));
|
|
request["stream"] = json!(stream);
|
|
let (status, body) = send(
|
|
byok::router(
|
|
store.clone(),
|
|
plugins.clone(),
|
|
Arc::new(fake_provider::FakeProvider::default()),
|
|
Some(byok::NativeForwarder::new(
|
|
store.clone(),
|
|
NetworkClients::new(store.clone()),
|
|
Duration::from_secs(10),
|
|
Duration::from_secs(10),
|
|
)),
|
|
),
|
|
"POST",
|
|
path,
|
|
Some("secret"),
|
|
request.clone(),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::OK, "{path} stream={stream}: {body}");
|
|
if stream {
|
|
assert_eq!(body, "data: {\"native_marker\":\"untouched-stream\"}\n\n");
|
|
} else {
|
|
assert_eq!(
|
|
serde_json::from_str::<Value>(&body).unwrap()["native_marker"],
|
|
"untouched-complete"
|
|
);
|
|
}
|
|
let (upstream_headers, forwarded) = receiver.recv().await.unwrap();
|
|
request["model"] = json!("native-model");
|
|
assert_eq!(forwarded, request);
|
|
if model_type == ModelType::Anthropic {
|
|
assert_eq!(upstream_headers["x-api-key"], "upstream");
|
|
assert_eq!(upstream_headers["anthropic-version"], "2023-06-01");
|
|
assert!(!upstream_headers.contains_key(header::AUTHORIZATION));
|
|
} else {
|
|
assert_eq!(upstream_headers[header::AUTHORIZATION], "Bearer upstream");
|
|
assert!(!upstream_headers.contains_key("x-api-key"));
|
|
}
|
|
}
|
|
}
|
|
let mut error_request = json!({"model":"chat/native-model","messages":[],"native_error":true});
|
|
let (status, body) = send(
|
|
byok::router(
|
|
store.clone(),
|
|
plugins,
|
|
Arc::new(fake_provider::FakeProvider::default()),
|
|
Some(byok::NativeForwarder::new(
|
|
store.clone(),
|
|
NetworkClients::new(store.clone()),
|
|
Duration::from_secs(10),
|
|
Duration::from_secs(10),
|
|
)),
|
|
),
|
|
"POST",
|
|
"/byok/v1/chat/completions",
|
|
Some("secret"),
|
|
error_request.clone(),
|
|
)
|
|
.await;
|
|
assert_eq!(status, StatusCode::TOO_MANY_REQUESTS);
|
|
assert_eq!(
|
|
serde_json::from_str::<Value>(&body).unwrap()["error"],
|
|
"native-rate-limit"
|
|
);
|
|
error_request["model"] = json!("native-model");
|
|
assert_eq!(receiver.recv().await.unwrap().1, error_request);
|
|
server.abort();
|
|
drop(directory);
|
|
}
|