1use std::sync::Arc;
4
5use axum::extract::State;
6use axum::http::HeaderMap;
7use axum::response::sse::{Event, Sse};
8use axum::response::{IntoResponse, Response};
9use axum::routing::{get, post};
10use axum::{Json, Router};
11use serde_json::json;
12
13use crate::error::GatewayError;
14use crate::gateway::{Gateway, GuardedStream, RequestCtx};
15use crate::routing::request::ChatRequest;
16use crate::routing::stream::{stream_item_to_sse_json, Accumulator, StreamItem};
17
18#[derive(Clone)]
19pub struct AppState {
20 pub gateway: Arc<Gateway>,
21}
22
23pub fn router(gateway: Arc<Gateway>) -> Router {
24 Router::new()
25 .route("/health", get(|| async { "ok" }))
26 .route("/v1/models", get(list_models))
27 .route("/v1/chat/completions", post(chat_completions))
28 .route("/v1/embeddings", post(embeddings))
29 .with_state(AppState { gateway })
30}
31
32fn request_ctx(headers: &HeaderMap) -> RequestCtx {
33 let header = |name: &str| {
34 headers
35 .get(name)
36 .and_then(|v| v.to_str().ok())
37 .map(str::to_string)
38 };
39 RequestCtx {
40 tenant: header("x-synapse-tenant"),
41 workspace: header("x-synapse-workspace"),
42 request_id: None,
43 }
44}
45
46async fn list_models(State(st): State<AppState>) -> impl IntoResponse {
47 let data = st
48 .gateway
49 .model_aliases()
50 .into_iter()
51 .map(|id| json!({ "id": id, "object": "model", "owned_by": "synapse" }))
52 .collect::<Vec<_>>();
53 Json(json!({ "object": "list", "data": data }))
54}
55
56async fn chat_completions(
57 State(st): State<AppState>,
58 headers: HeaderMap,
59 Json(req): Json<ChatRequest>,
60) -> Result<Response, GatewayError> {
61 let request_id = uuid::Uuid::new_v4().to_string();
62 let ctx = RequestCtx {
64 request_id: Some(request_id.clone()),
65 ..request_ctx(&headers)
66 };
67
68 if req.stream == Some(true) {
69 let stream = st.gateway.chat_stream(req, &ctx).await?;
70 return Ok(Sse::new(sse_body(stream, request_id)).into_response());
71 }
72
73 let completion = st.gateway.chat(req, &ctx).await?;
74 Ok(Json(openai_json(&completion, &request_id)).into_response())
75}
76
77async fn embeddings(
78 State(st): State<AppState>,
79 headers: HeaderMap,
80 Json(req): Json<crate::embeddings::EmbeddingRequest>,
81) -> Result<Response, GatewayError> {
82 let ctx = request_ctx(&headers);
83 let resp = st.gateway.embed(req, ctx).await?;
84 Ok(Json(resp).into_response())
85}
86
87fn openai_json(c: &crate::routing::executor::Completion, request_id: &str) -> serde_json::Value {
90 let mut acc = Accumulator::default();
91 if c.tool_calls.is_empty() {
92 acc.push(StreamItem::Delta(c.content.clone()));
93 } else {
94 for (i, tc) in c.tool_calls.iter().enumerate() {
95 acc.push(StreamItem::ToolCallDelta {
96 index: i as u32,
97 id: Some(tc.id.clone()),
98 name: Some(tc.name.clone()),
99 args_fragment: tc.arguments.clone(),
100 });
101 }
102 }
103 acc.push(StreamItem::Done {
104 input_tokens: c.input_tokens,
105 output_tokens: c.output_tokens,
106 finish_reason: c.finish_reason,
107 });
108 acc.to_openai_response(request_id, &c.model)
109}
110
111fn sse_body(
113 stream: GuardedStream,
114 request_id: String,
115) -> impl futures::Stream<Item = Result<Event, std::convert::Infallible>> {
116 use futures::StreamExt;
117 let model = stream.model().to_string();
118 stream
119 .map(move |item| match item {
120 Ok(it) => {
121 let json = stream_item_to_sse_json(&it, &request_id, &model);
122 Ok(Event::default().data(json.to_string()))
123 }
124 Err(e) => {
125 let err = json!({
126 "error": { "type": "upstream_error", "message": e.to_string(), "code": "upstream_error" }
127 });
128 Ok(Event::default().data(err.to_string()))
129 }
130 })
131 .chain(futures::stream::once(async { Ok(Event::default().data("[DONE]")) }))
132}