diff --git a/bin/pudox-run/src/main.rs b/bin/pudox-run/src/main.rs index 4cb6566..98083e5 100644 --- a/bin/pudox-run/src/main.rs +++ b/bin/pudox-run/src/main.rs @@ -8,9 +8,17 @@ use async_openai::types::chat::{ }; use async_openai::Client; -const PUDOXD_TASK_URL: &str = "http://localhost:7790/task"; -const PUDOXD_RESULT_URL: &str = "http://localhost:7790/task/result"; -const PUDOXD_STREAM_URL: &str = "http://localhost:7790/task/stream"; +macro_rules! pudoxd_url { + ($path:literal) => { + concat!("http://localhost:7790", $path) + }; +} + +const PUDOXD_OPENAI_URL: &str = pudoxd_url!("/v1"); +const PUDOXD_TASK_URL: &str = pudoxd_url!("/task"); +const PUDOXD_RESULT_URL: &str = pudoxd_url!("/task/result"); +const PUDOXD_STREAM_URL: &str = pudoxd_url!("/task/stream"); +const PUDOXD_INFERENCE_URL: &str = pudoxd_url!("/v1/chat/completions"); #[tokio::main] async fn main() -> Result<(), Box> { @@ -123,7 +131,7 @@ async fn process_task( ) -> Result> { let client = Client::with_config( OpenAIConfig::new() - .with_api_base("http://ai") + .with_api_base(PUDOXD_OPENAI_URL) .with_api_key("none"), ); @@ -293,7 +301,7 @@ async fn summarize_session( }); let resp: serde_json::Value = reqwest::Client::new() - .post("http://ai/v1/chat/completions") + .post(PUDOXD_INFERENCE_URL) .json(&body) .send() .await? @@ -342,7 +350,7 @@ async fn generate_pull_request_info( for attempt in 1..=3u32 { tracing::info!("generate_pull_request_info: attempt {}/3", attempt); let resp: serde_json::Value = reqwest::Client::new() - .post("http://ai/v1/chat/completions") + .post(PUDOXD_INFERENCE_URL) .json(&body) .send() .await? diff --git a/bin/pudoxd/Cargo.toml b/bin/pudoxd/Cargo.toml index 2db1869..25fa47e 100644 --- a/bin/pudoxd/Cargo.toml +++ b/bin/pudoxd/Cargo.toml @@ -12,7 +12,7 @@ axum = "0.7" tokio = { version = "1", features = ["full"] } tracing = "0.1.44" tracing-subscriber = "0.3.23" -reqwest = { version = "0.12", features = ["json"] } +reqwest = { version = "0.12", features = ["json", "stream"] } serde = { version = "1.0", features = ["derive"] } serde_json = "1.0" git2 = { version = "0.20", default-features = false, features = ["vendored-libgit2", "vendored-openssl", "https"] } diff --git a/bin/pudoxd/src/main.rs b/bin/pudoxd/src/main.rs index 3d3fc9e..769d70e 100644 --- a/bin/pudoxd/src/main.rs +++ b/bin/pudoxd/src/main.rs @@ -1,6 +1,7 @@ use axum::{ + body::Body, extract::State, - http::StatusCode, + http::{Request, Response, StatusCode}, response::{ sse::{Event, KeepAlive, Sse}, IntoResponse, @@ -25,6 +26,7 @@ use crate::session::{run_session, Session, SessionMessage}; struct App { linear: Linear, sessions: Arc>>, + http: reqwest::Client, } #[tokio::main] @@ -36,6 +38,7 @@ async fn main() -> Result<(), Box> { let app = App { linear: Linear::new(), sessions: Arc::new(RwLock::new(HashMap::new())), + http: reqwest::Client::new(), }; let router = Router::new() @@ -43,6 +46,7 @@ async fn main() -> Result<(), Box> { .route("/task", get(task_handler)) .route("/task/result", post(task_result_handler)) .route("/task/stream", get(stream_handler)) + .route("/v1/chat/completions", post(chat_completions_proxy)) .with_state(app); let addr = SocketAddr::from(([0, 0, 0, 0], 7790)); @@ -344,6 +348,57 @@ async fn handle_task_result( } } +async fn chat_completions_proxy(State(app): State, req: Request) -> impl IntoResponse { + let content_type = req + .headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + .unwrap_or("application/json") + .to_string(); + + let body_bytes = match axum::body::to_bytes(req.into_body(), usize::MAX).await { + Ok(b) => b, + Err(_) => { + return Response::builder() + .status(StatusCode::BAD_GATEWAY) + .body(Body::empty()) + .unwrap(); + } + }; + + let upstream = match app + .http + .post("http://ai/v1/chat/completions") + .header("content-type", content_type) + .body(body_bytes) + .send() + .await + { + Ok(r) => r, + Err(e) => { + tracing::error!("Upstream inference request failed: {}", e); + return Response::builder() + .status(StatusCode::BAD_GATEWAY) + .body(Body::empty()) + .unwrap(); + } + }; + + let status = upstream.status(); + let resp_content_type = upstream + .headers() + .get("content-type") + .and_then(|v| v.to_str().ok()) + .unwrap_or("application/json") + .to_string(); + + Response::builder() + .status(status) + .header("content-type", resp_content_type) + .body(Body::from_stream(upstream.bytes_stream())) + .unwrap() +} + async fn stream_handler( State(app): State, headers: axum::http::HeaderMap,