Skip to main content

weft_core/rpc/
mod.rs

1use crate::api::openai_compat::AppState;
2use crate::config::AppConfig;
3use anyhow::{anyhow, Result};
4use serde::{Deserialize, Serialize};
5use serde_json::{json, Value};
6use std::io::{self, BufRead, Write};
7
8#[derive(Debug, Deserialize)]
9pub struct JsonRpcRequest {
10    pub jsonrpc: String,
11    pub id: Option<Value>,
12    pub method: String,
13    #[serde(default)]
14    pub params: Option<Value>,
15}
16
17#[derive(Debug, Serialize, PartialEq)]
18pub struct JsonRpcResponse {
19    pub jsonrpc: &'static str,
20    #[serde(skip_serializing_if = "Option::is_none")]
21    pub id: Option<Value>,
22    #[serde(skip_serializing_if = "Option::is_none")]
23    pub result: Option<Value>,
24    #[serde(skip_serializing_if = "Option::is_none")]
25    pub error: Option<JsonRpcError>,
26}
27
28#[derive(Debug, Serialize, PartialEq)]
29pub struct JsonRpcError {
30    pub code: i64,
31    pub message: String,
32}
33
34pub async fn handle_request(request: JsonRpcRequest, state: Option<&AppState>) -> JsonRpcResponse {
35    if request.jsonrpc != "2.0" {
36        return error_response(request.id, -32600, "Invalid Request");
37    }
38
39    match request.method.as_str() {
40        "health" => success_response(request.id, health_result()),
41        "models/list" => {
42            if let Some(state) = state {
43                let config = state.config.read().await;
44                success_response(request.id, model_list_result(&config))
45            } else {
46                success_response(request.id, model_list_from_providers(&[]))
47            }
48        }
49        _ => error_response(request.id, -32601, "Method not found"),
50    }
51}
52
53pub fn health_result() -> Value {
54    json!({
55        "status": "ok",
56        "version": env!("CARGO_PKG_VERSION"),
57    })
58}
59
60pub fn model_list_result(config: &AppConfig) -> Value {
61    let providers: Vec<(&str, &[String])> = config
62        .providers
63        .iter()
64        .map(|provider| (provider.name.as_str(), provider.models.as_slice()))
65        .collect();
66    model_list_from_providers(&providers)
67}
68
69fn model_list_from_providers(providers: &[(&str, &[String])]) -> Value {
70    let models: Vec<Value> = providers
71        .iter()
72        .flat_map(|(provider, models)| {
73            models.iter().map(move |model| {
74                json!({
75                    "id": model,
76                    "object": "model",
77                    "owned_by": provider,
78                })
79            })
80        })
81        .collect();
82
83    json!({
84        "object": "list",
85        "data": models,
86    })
87}
88
89pub async fn serve_stdio(state: Option<AppState>) -> Result<()> {
90    let stdin = io::stdin();
91    let mut stdout = io::stdout();
92
93    for line in stdin.lock().lines() {
94        let line = line?;
95        if line.trim().is_empty() {
96            continue;
97        }
98
99        let response = match serde_json::from_str::<JsonRpcRequest>(&line) {
100            Ok(request) => handle_request(request, state.as_ref()).await,
101            Err(_) => error_response(None, -32700, "Parse error"),
102        };
103        serde_json::to_writer(&mut stdout, &response)?;
104        stdout.write_all(b"\n")?;
105        stdout.flush()?;
106    }
107
108    Ok(())
109}
110
111pub fn parse_single_request(input: &str) -> Result<JsonRpcRequest> {
112    serde_json::from_str(input).map_err(|error| anyhow!("invalid JSON-RPC request: {error}"))
113}
114
115fn success_response(id: Option<Value>, result: Value) -> JsonRpcResponse {
116    JsonRpcResponse {
117        jsonrpc: "2.0",
118        id,
119        result: Some(result),
120        error: None,
121    }
122}
123
124fn error_response(id: Option<Value>, code: i64, message: &str) -> JsonRpcResponse {
125    JsonRpcResponse {
126        jsonrpc: "2.0",
127        id,
128        result: None,
129        error: Some(JsonRpcError {
130            code,
131            message: message.to_string(),
132        }),
133    }
134}
135
136#[cfg(test)]
137mod tests {
138    use super::*;
139
140    #[tokio::test]
141    async fn handles_health_request() {
142        let request = parse_single_request(r#"{"jsonrpc":"2.0","id":1,"method":"health"}"#)
143            .expect("valid request");
144
145        let response = handle_request(request, None).await;
146
147        assert_eq!(response.id, Some(json!(1)));
148        assert_eq!(response.error, None);
149        assert_eq!(response.result.unwrap()["status"], "ok");
150    }
151
152    #[tokio::test]
153    async fn handles_model_list_without_state() {
154        let request =
155            parse_single_request(r#"{"jsonrpc":"2.0","id":"models","method":"models/list"}"#)
156                .expect("valid request");
157
158        let response = handle_request(request, None).await;
159
160        assert_eq!(response.id, Some(json!("models")));
161        assert_eq!(response.error, None);
162        assert_eq!(response.result.unwrap(), json!({"object":"list","data":[]}));
163    }
164
165    #[tokio::test]
166    async fn rejects_unknown_method() {
167        let request = parse_single_request(r#"{"jsonrpc":"2.0","id":2,"method":"missing"}"#)
168            .expect("valid request");
169
170        let response = handle_request(request, None).await;
171
172        assert_eq!(
173            response.error,
174            Some(JsonRpcError {
175                code: -32601,
176                message: "Method not found".to_string(),
177            })
178        );
179    }
180}