use std::future::Future;
use std::pin::Pin;
use std::time::Duration;
use serde_json::{json, Value};
pub type MediaResponse = (u16, Value);
pub trait ImageHost: Send + Sync {
fn sd_base_url(&self) -> String;
fn gateway_url(&self) -> String;
fn gateway_token(&self) -> Option<String>;
fn start_local_engine(&self) -> Pin<Box<dyn Future<Output = Result<(), String>> + Send + '_>>;
}
pub const CLOUD_PROVIDERS: [&str; 3] = ["openrouter", "replicate", "fal"];
pub fn cloud_provider(body: &Value) -> Option<String> {
body.get("provider")
.and_then(Value::as_str)
.map(|s| s.trim().to_lowercase())
.filter(|s| CLOUD_PROVIDERS.contains(&s.as_str()))
}
pub fn media_client() -> reqwest::Client {
reqwest::Client::builder()
.user_agent("ryu-core/0.1")
.timeout(Duration::from_secs(600))
.build()
.expect("reqwest client")
}
pub async fn forward_to_gateway(
host: &impl ImageHost,
modality: &str,
endpoint: &str,
provider: &str,
body: Value,
) -> MediaResponse {
let base = host.gateway_url();
let url = format!("{}{endpoint}", base.trim_end_matches('/'));
let slot_header = format!("x-ryu-slot-{modality}-provider");
let mut req = media_client()
.post(&url)
.header(slot_header, provider)
.json(&body);
if let Some(t) = host.gateway_token() {
req = req.bearer_auth(t);
}
let resp = match req.send().await {
Ok(r) => r,
Err(e) => {
return (
502,
json!({
"error": format!("cloud media gateway not reachable at {url}: {e}")
}),
);
}
};
let status = resp.status();
let bytes = resp.bytes().await.unwrap_or_default();
let value: Value = serde_json::from_slice(&bytes)
.unwrap_or_else(|_| json!({ "raw": String::from_utf8_lossy(&bytes) }));
if !status.is_success() {
return (
502,
json!({ "error": format!("cloud media provider returned {status}"), "detail": value }),
);
}
(200, value)
}
pub async fn proxy(base_url: &str, endpoint: &str, body: Value) -> MediaResponse {
let url = format!("{base_url}{endpoint}");
let resp = match media_client().post(&url).json(&body).send().await {
Ok(r) => r,
Err(e) => {
return (
502,
json!({
"error": format!(
"stable-diffusion.cpp media engine not reachable at {url}: {e}. \
Install + start `sdcpp` from the Store first."
)
}),
);
}
};
let status = resp.status();
let bytes = resp.bytes().await.unwrap_or_default();
let value: Value = serde_json::from_slice(&bytes)
.unwrap_or_else(|_| json!({ "raw": String::from_utf8_lossy(&bytes) }));
if !status.is_success() {
return (
502,
json!({ "error": format!("media engine returned {status}"), "detail": value }),
);
}
(200, value)
}
pub async fn generate(host: &impl ImageHost, mut body: Value) -> MediaResponse {
if body
.get("prompt")
.and_then(Value::as_str)
.unwrap_or("")
.trim()
.is_empty()
{
return (
400,
json!({ "error": "missing `prompt` (the text to render)" }),
);
}
if let Some(obj) = body.as_object_mut() {
obj.entry("n").or_insert(json!(1));
}
if let Some(provider) = cloud_provider(&body) {
return forward_to_gateway(host, "image", "/v1/images/generations", &provider, body).await;
}
if let Err(e) = host.start_local_engine().await {
tracing::debug!("sdcpp lazy start skipped: {e:#}");
}
proxy(&host.sd_base_url(), "/v1/images/generations", body).await
}
#[cfg(test)]
mod tests {
use super::*;
struct FakeHost;
impl ImageHost for FakeHost {
fn sd_base_url(&self) -> String {
"http://127.0.0.1:8083".into()
}
fn gateway_url(&self) -> String {
"http://127.0.0.1:7981".into()
}
fn gateway_token(&self) -> Option<String> {
None
}
fn start_local_engine(
&self,
) -> Pin<Box<dyn Future<Output = Result<(), String>> + Send + '_>> {
Box::pin(async { Ok(()) })
}
}
#[test]
fn cloud_provider_selects_known_and_normalizes() {
assert_eq!(
cloud_provider(&json!({ "provider": " Replicate " })),
Some("replicate".into())
);
assert_eq!(
cloud_provider(&json!({ "provider": "fal" })),
Some("fal".into())
);
}
#[test]
fn cloud_provider_rejects_unknown_or_absent() {
assert_eq!(cloud_provider(&json!({ "provider": "midjourney" })), None);
assert_eq!(cloud_provider(&json!({ "prompt": "a cat" })), None);
}
#[tokio::test]
async fn generate_rejects_empty_prompt() {
let (code, body) = generate(&FakeHost, json!({ "prompt": " " })).await;
assert_eq!(code, 400);
assert!(body.get("error").is_some());
}
#[tokio::test]
async fn generate_local_unreachable_engine_is_bad_gateway() {
let (code, body) = generate(&FakeHost, json!({ "prompt": "a corgi" })).await;
assert_eq!(code, 502);
assert!(body["error"]
.as_str()
.unwrap_or("")
.contains("not reachable"));
}
}