feat(web): image search and generate endpoints on the canvas server
This commit is contained in:
parent
2314b736b2
commit
2ff78e7d44
|
|
@ -69,5 +69,7 @@ pub mod web_canvas_server;
|
|||
pub mod web_chat_standard;
|
||||
pub mod web_credential_policy;
|
||||
pub mod web_credentials;
|
||||
mod web_image_generate;
|
||||
mod web_image_search;
|
||||
pub mod web_static;
|
||||
pub mod zode_import;
|
||||
|
|
|
|||
|
|
@ -1786,6 +1786,68 @@ fn serve_one<S: Read + Write>(
|
|||
.map_err(|e| format!("ai standard: {e}"))
|
||||
.map(|()| false);
|
||||
}
|
||||
// Image panel Search popover (desktop `image_panel_host` parity). Long
|
||||
// blocking network (8 s timeout × ladder), so it runs on this
|
||||
// connection's own thread AFTER the brief parse-under-lock — the REST
|
||||
// handler below holds the state lock for its whole body and must not
|
||||
// host provider dials. Living under `/api/ai/` keeps it inside the
|
||||
// sensitive-POST origin gate and the managed-mode token gate.
|
||||
if req.method == "POST" && req.path == "/api/ai/image/search" {
|
||||
let parsed = {
|
||||
let guard = state.lock().unwrap_or_else(|p| p.into_inner());
|
||||
crate::web_image_search::parse_search_request(&req.body, &guard.editor)
|
||||
};
|
||||
let (status, body) = match parsed {
|
||||
Ok((query, credentials)) => {
|
||||
let outcome =
|
||||
crate::web_image_search::run_search_blocking(&query, credentials.as_ref());
|
||||
(
|
||||
"200 OK",
|
||||
crate::web_image_search::search_outcome_to_json(&outcome),
|
||||
)
|
||||
}
|
||||
Err(message) => (
|
||||
"400 Bad Request",
|
||||
serde_json::json!({ "ok": false, "error": message }).to_string(),
|
||||
),
|
||||
};
|
||||
return crate::mcp_serve::write_mcp_http_response_with_origin(
|
||||
stream,
|
||||
status,
|
||||
&body,
|
||||
cors_origin,
|
||||
)
|
||||
.map(|()| false);
|
||||
}
|
||||
// Image panel Generate popover (desktop `image_generate_host` parity).
|
||||
// Same threading rules as the search route; Replicate polling can run
|
||||
// for minutes.
|
||||
if req.method == "POST" && req.path == "/api/ai/image/generate" {
|
||||
let parsed = {
|
||||
let guard = state.lock().unwrap_or_else(|p| p.into_inner());
|
||||
crate::web_image_generate::parse_generate_request(&req.body, &guard.editor)
|
||||
};
|
||||
let (status, body) = match parsed {
|
||||
Ok(request) => match crate::web_image_generate::run_generate_blocking(&request) {
|
||||
Ok(url) => ("200 OK", crate::web_image_generate::generate_ok_json(&url)),
|
||||
Err(message) => (
|
||||
"502 Bad Gateway",
|
||||
crate::web_image_generate::generate_error_json(&message),
|
||||
),
|
||||
},
|
||||
Err(message) => (
|
||||
"400 Bad Request",
|
||||
crate::web_image_generate::generate_error_json(&message),
|
||||
),
|
||||
};
|
||||
return crate::mcp_serve::write_mcp_http_response_with_origin(
|
||||
stream,
|
||||
status,
|
||||
&body,
|
||||
cors_origin,
|
||||
)
|
||||
.map(|()| false);
|
||||
}
|
||||
// All `/api/mcp/*` REST paths go to the REST handler — including ones this
|
||||
// daemon doesn't implement yet, which it answers with 404 rather than
|
||||
// mis-routing them into the JSON-RPC dispatch below.
|
||||
|
|
|
|||
615
crates/op-host-services/src/web_image_generate.rs
Normal file
615
crates/op-host-services/src/web_image_generate.rs
Normal file
|
|
@ -0,0 +1,615 @@
|
|||
//! Browser image-generation route for the web daemon.
|
||||
//!
|
||||
//! `POST /api/ai/image/generate` mirrors the desktop Generate popover
|
||||
//! backend (`op-host-desktop/src/image_generate_host.rs`: OpenAI /
|
||||
//! OpenAI-compatible, Gemini inline-image, Replicate prediction polling).
|
||||
//! The generation profile comes from the request body (browser-held keys)
|
||||
//! or falls back to the daemon's persisted agent settings.
|
||||
//!
|
||||
//! Endpoint trust follows the chat route's rules: provider DEFAULT hosts
|
||||
//! are product constants and dial with a plain client, but any custom
|
||||
//! `base_url` reaching this browser-facing route is screened through
|
||||
//! `web_credentials` + `provider_dial` (reserved-address rejection +
|
||||
//! connect-time pinning) unless explicitly allowlisted. The provider's
|
||||
//! RESULT url is screened the same way before download — a custom
|
||||
//! provider's response is browser-controlled data and must not become an
|
||||
//! SSRF hop through the daemon.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use op_editor_core::agent_settings::{ImageGenProfile, ImageGenProvider, ImageTestStatus};
|
||||
|
||||
use crate::provider_dial::EndpointDialPolicy;
|
||||
use crate::web_image_search::fetch_image_data_url;
|
||||
|
||||
/// Truncation cap for surfaced provider errors (TS parity).
|
||||
const ERROR_MESSAGE_CAP: usize = 200;
|
||||
|
||||
pub(crate) struct WebImageGenerateRequest {
|
||||
pub(crate) prompt: String,
|
||||
pub(crate) width: Option<f64>,
|
||||
pub(crate) height: Option<f64>,
|
||||
pub(crate) profile: ImageGenProfile,
|
||||
/// Whether `profile` carries a custom endpoint that must be screened.
|
||||
pub(crate) custom_endpoint: bool,
|
||||
}
|
||||
|
||||
/// Parse the request body, taking the browser-supplied profile when present
|
||||
/// and falling back to the daemon's persisted active profile.
|
||||
pub(crate) fn parse_generate_request(
|
||||
body: &str,
|
||||
state: &op_editor_core::EditorState,
|
||||
) -> Result<WebImageGenerateRequest, String> {
|
||||
let value: serde_json::Value =
|
||||
serde_json::from_str(body).map_err(|_| "invalid request body".to_string())?;
|
||||
let obj = value.as_object().ok_or("invalid request body")?;
|
||||
let prompt = obj
|
||||
.get("prompt")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|p| !p.is_empty())
|
||||
.ok_or("missing prompt")?;
|
||||
let width = obj.get("width").and_then(serde_json::Value::as_f64);
|
||||
let height = obj.get("height").and_then(serde_json::Value::as_f64);
|
||||
let profile = match obj.get("profile").and_then(serde_json::Value::as_object) {
|
||||
Some(profile) => parse_profile(profile)?,
|
||||
None => daemon_active_profile(state).ok_or("image generation not configured")?,
|
||||
};
|
||||
let custom_endpoint = profile
|
||||
.base_url
|
||||
.as_deref()
|
||||
.is_some_and(|base| !base.trim().is_empty());
|
||||
Ok(WebImageGenerateRequest {
|
||||
prompt: prompt.to_string(),
|
||||
width,
|
||||
height,
|
||||
profile,
|
||||
custom_endpoint,
|
||||
})
|
||||
}
|
||||
|
||||
fn parse_profile(
|
||||
profile: &serde_json::Map<String, serde_json::Value>,
|
||||
) -> Result<ImageGenProfile, String> {
|
||||
let provider = match profile
|
||||
.get("provider")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
{
|
||||
"openai" => ImageGenProvider::OpenAi,
|
||||
"gemini" => ImageGenProvider::Gemini,
|
||||
"replicate" => ImageGenProvider::Replicate,
|
||||
"custom" => ImageGenProvider::Custom,
|
||||
_ => return Err("unknown image provider".to_string()),
|
||||
};
|
||||
let api_key = profile
|
||||
.get("api_key")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|k| !k.is_empty())
|
||||
.ok_or("missing api key")?;
|
||||
let model = profile
|
||||
.get("model")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|m| !m.is_empty())
|
||||
.ok_or("missing model")?;
|
||||
let base_url = profile
|
||||
.get("base_url")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|b| !b.is_empty())
|
||||
.map(str::to_string);
|
||||
Ok(ImageGenProfile {
|
||||
id: String::new(),
|
||||
name: String::new(),
|
||||
provider,
|
||||
api_key: api_key.to_string(),
|
||||
model: model.to_string(),
|
||||
base_url,
|
||||
test_status: ImageTestStatus::Idle,
|
||||
})
|
||||
}
|
||||
|
||||
/// The daemon's persisted active profile (desktop `active_image_gen_profile`).
|
||||
fn daemon_active_profile(state: &op_editor_core::EditorState) -> Option<ImageGenProfile> {
|
||||
let settings = &state.editor_ui.agent_settings;
|
||||
settings
|
||||
.image_gen_profiles
|
||||
.iter()
|
||||
.find(|p| Some(&p.id) == settings.active_image_gen_profile_id.as_ref())
|
||||
.or_else(|| settings.image_gen_profiles.first())
|
||||
.filter(|p| !p.api_key.trim().is_empty())
|
||||
.cloned()
|
||||
}
|
||||
|
||||
/// Run one generation on the calling thread (the connection's own thread —
|
||||
/// the caller must NOT hold the state lock). Returns the final `data:` URL.
|
||||
pub(crate) fn run_generate_blocking(request: &WebImageGenerateRequest) -> Result<String, String> {
|
||||
let runtime = tokio::runtime::Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
.map_err(|e| format!("tokio runtime: {e}"))?;
|
||||
runtime.block_on(run_generate(request))
|
||||
}
|
||||
|
||||
async fn run_generate(request: &WebImageGenerateRequest) -> Result<String, String> {
|
||||
let profile = &request.profile;
|
||||
// Screen browser-supplied custom endpoints before any dial; default
|
||||
// provider hosts are product constants and stay trusted.
|
||||
let endpoint_policy = if request.custom_endpoint {
|
||||
let base = profile.base_url.as_deref().unwrap_or_default();
|
||||
crate::web_credentials::validate_web_provider_base_url(base)
|
||||
.map_err(|_| "provider endpoint is not allowed".to_string())?;
|
||||
let allowlist = std::env::var(crate::web_credentials::WEB_AI_ENDPOINT_ALLOWLIST_ENV).ok();
|
||||
crate::provider_dial::web_dial_policy_for(base, allowlist.as_deref())
|
||||
} else {
|
||||
EndpointDialPolicy::Trusted
|
||||
};
|
||||
let client = match endpoint_policy {
|
||||
EndpointDialPolicy::Trusted => generate_client()?,
|
||||
EndpointDialPolicy::PublicOnly => {
|
||||
let base = profile.base_url.as_deref().unwrap_or_default();
|
||||
crate::provider_dial::client_for(endpoint_policy, base).await?
|
||||
}
|
||||
};
|
||||
let url = match profile.provider {
|
||||
ImageGenProvider::OpenAi | ImageGenProvider::Custom => {
|
||||
generate_openai(
|
||||
&client,
|
||||
&request.prompt,
|
||||
profile,
|
||||
request.width,
|
||||
request.height,
|
||||
)
|
||||
.await?
|
||||
}
|
||||
ImageGenProvider::Gemini => {
|
||||
generate_gemini(
|
||||
&client,
|
||||
&request.prompt,
|
||||
profile,
|
||||
request.width,
|
||||
request.height,
|
||||
)
|
||||
.await?
|
||||
}
|
||||
ImageGenProvider::Replicate => {
|
||||
generate_replicate(
|
||||
&client,
|
||||
&request.prompt,
|
||||
profile,
|
||||
request.width,
|
||||
request.height,
|
||||
)
|
||||
.await?
|
||||
}
|
||||
};
|
||||
if url.starts_with("data:") {
|
||||
// Inline base64 (Gemini / OpenAI b64_json). Unlike the desktop, no
|
||||
// down-scale pass here — image_downscale needs skia and this crate
|
||||
// stays GL-free.
|
||||
return Ok(url);
|
||||
}
|
||||
// Remote URL → embed as a data URL so the preview paints and the applied
|
||||
// src stays renderable offline. Any CUSTOM provider's result URL is
|
||||
// browser/endpoint-controlled data and is screened independently — an
|
||||
// allowlist entry trusts the configured endpoint, not every URL its
|
||||
// responses point at (the download host can be allowlisted separately).
|
||||
// Default provider hosts keep the trusted client for their own CDNs.
|
||||
let download_client = if request.custom_endpoint {
|
||||
let allowlist = std::env::var(crate::web_credentials::WEB_AI_ENDPOINT_ALLOWLIST_ENV).ok();
|
||||
let policy = crate::provider_dial::web_dial_policy_for(&url, allowlist.as_deref());
|
||||
crate::provider_dial::client_for(policy, &url)
|
||||
.await
|
||||
.map_err(|_| "generated image URL is not allowed".to_string())?
|
||||
} else {
|
||||
client.clone()
|
||||
};
|
||||
fetch_image_data_url(&download_client, &url)
|
||||
.await
|
||||
.ok_or_else(|| "generated image could not be downloaded".to_string())
|
||||
}
|
||||
|
||||
fn generate_client() -> Result<reqwest::Client, String> {
|
||||
reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(120))
|
||||
.user_agent(concat!("openpencil-web-daemon/", env!("CARGO_PKG_VERSION")))
|
||||
.build()
|
||||
.map_err(|e| format!("http client: {e}"))
|
||||
}
|
||||
|
||||
/// Cap + JSON-wrap an error for the reply body.
|
||||
pub(crate) fn generate_error_json(message: &str) -> String {
|
||||
let message: String = message.chars().take(ERROR_MESSAGE_CAP).collect();
|
||||
serde_json::json!({ "ok": false, "error": message }).to_string()
|
||||
}
|
||||
|
||||
pub(crate) fn generate_ok_json(url: &str) -> String {
|
||||
serde_json::json!({ "ok": true, "url": url }).to_string()
|
||||
}
|
||||
|
||||
/// TS `mapToOpenAISize`.
|
||||
fn openai_size(width: Option<f64>, height: Option<f64>) -> &'static str {
|
||||
let (Some(w), Some(h)) = (width, height) else {
|
||||
return "1024x1024";
|
||||
};
|
||||
let ratio = w / h;
|
||||
if ratio > 1.3 {
|
||||
"1792x1024"
|
||||
} else if ratio < 0.77 {
|
||||
"1024x1792"
|
||||
} else {
|
||||
"1024x1024"
|
||||
}
|
||||
}
|
||||
|
||||
/// TS `mapToGeminiAspectRatio`.
|
||||
fn gemini_aspect_ratio(width: Option<f64>, height: Option<f64>) -> Option<&'static str> {
|
||||
let (Some(w), Some(h)) = (width, height) else {
|
||||
return None;
|
||||
};
|
||||
let ratio = w / h;
|
||||
Some(if ratio > 1.6 {
|
||||
"16:9"
|
||||
} else if ratio > 1.3 {
|
||||
"4:3"
|
||||
} else if ratio < 0.625 {
|
||||
"9:16"
|
||||
} else if ratio < 0.77 {
|
||||
"3:4"
|
||||
} else {
|
||||
"1:1"
|
||||
})
|
||||
}
|
||||
|
||||
fn provider_error(provider: &str, status: reqwest::StatusCode, body: &str) -> String {
|
||||
// TS: prefer the provider's error.message, else status + slice.
|
||||
if let Ok(json) = serde_json::from_str::<serde_json::Value>(body) {
|
||||
if let Some(message) = json
|
||||
.get("error")
|
||||
.and_then(|e| e.get("message"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| json.get("detail").and_then(serde_json::Value::as_str))
|
||||
{
|
||||
return message.chars().take(ERROR_MESSAGE_CAP).collect();
|
||||
}
|
||||
}
|
||||
let mut msg = format!("{provider} returned {}", status.as_u16());
|
||||
if !body.is_empty() {
|
||||
msg.push_str(": ");
|
||||
msg.push_str(&body.chars().take(150).collect::<String>());
|
||||
}
|
||||
msg
|
||||
}
|
||||
|
||||
async fn generate_openai(
|
||||
client: &reqwest::Client,
|
||||
prompt: &str,
|
||||
profile: &ImageGenProfile,
|
||||
width: Option<f64>,
|
||||
height: Option<f64>,
|
||||
) -> Result<String, String> {
|
||||
let base = profile
|
||||
.base_url
|
||||
.as_deref()
|
||||
.filter(|b| !b.trim().is_empty())
|
||||
.unwrap_or("https://api.openai.com");
|
||||
let endpoint = format!("{}/v1/images/generations", base.trim_end_matches('/'));
|
||||
let resp = client
|
||||
.post(endpoint)
|
||||
.bearer_auth(profile.api_key.trim())
|
||||
.json(&serde_json::json!({
|
||||
"model": profile.model,
|
||||
"prompt": prompt,
|
||||
"n": 1,
|
||||
"size": openai_size(width, height),
|
||||
"response_format": "url",
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("OpenAI request failed: {e}"))?;
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
if !status.is_success() {
|
||||
return Err(provider_error("OpenAI", status, &body));
|
||||
}
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&body).map_err(|e| format!("OpenAI response parse: {e}"))?;
|
||||
json.get("data")
|
||||
.and_then(|d| d.as_array())
|
||||
.and_then(|d| d.first())
|
||||
.and_then(|d| {
|
||||
d.get("url")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_string)
|
||||
.or_else(|| {
|
||||
// Some OpenAI-compatible providers return b64_json.
|
||||
d.get("b64_json")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(|b64| format!("data:image/png;base64,{b64}"))
|
||||
})
|
||||
})
|
||||
.ok_or_else(|| "OpenAI response missing image URL".to_string())
|
||||
}
|
||||
|
||||
async fn generate_gemini(
|
||||
client: &reqwest::Client,
|
||||
prompt: &str,
|
||||
profile: &ImageGenProfile,
|
||||
width: Option<f64>,
|
||||
height: Option<f64>,
|
||||
) -> Result<String, String> {
|
||||
let base = profile
|
||||
.base_url
|
||||
.as_deref()
|
||||
.filter(|b| !b.trim().is_empty())
|
||||
.unwrap_or("https://generativelanguage.googleapis.com");
|
||||
let endpoint = format!(
|
||||
"{}/v1beta/models/{}:generateContent?key={}",
|
||||
base.trim_end_matches('/'),
|
||||
profile.model,
|
||||
profile.api_key.trim()
|
||||
);
|
||||
let mut generation_config = serde_json::json!({
|
||||
"responseModalities": ["TEXT", "IMAGE"],
|
||||
});
|
||||
if let Some(aspect) = gemini_aspect_ratio(width, height) {
|
||||
generation_config["imageConfig"] = serde_json::json!({ "aspectRatio": aspect });
|
||||
}
|
||||
let resp = client
|
||||
.post(endpoint)
|
||||
.json(&serde_json::json!({
|
||||
"contents": [{ "parts": [{ "text": prompt }] }],
|
||||
"generationConfig": generation_config,
|
||||
}))
|
||||
.send()
|
||||
.await
|
||||
// `without_url`: the Gemini endpoint carries `?key=…`, and a plain
|
||||
// reqwest error Display would echo that URL — API key included —
|
||||
// back to the browser.
|
||||
.map_err(|e| format!("Gemini request failed: {}", e.without_url()))?;
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
if !status.is_success() {
|
||||
return Err(provider_error("Gemini", status, &body));
|
||||
}
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&body).map_err(|e| format!("Gemini response parse: {e}"))?;
|
||||
let parts = json
|
||||
.get("candidates")
|
||||
.and_then(|c| c.as_array())
|
||||
.and_then(|c| c.first())
|
||||
.and_then(|c| c.get("content"))
|
||||
.and_then(|c| c.get("parts"))
|
||||
.and_then(|p| p.as_array());
|
||||
let Some(parts) = parts else {
|
||||
return Err("Gemini response missing inline image data".to_string());
|
||||
};
|
||||
for part in parts {
|
||||
let Some(inline) = part.get("inlineData") else {
|
||||
continue;
|
||||
};
|
||||
let mime = inline
|
||||
.get("mimeType")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("");
|
||||
if !mime.starts_with("image/") {
|
||||
continue;
|
||||
}
|
||||
if let Some(data) = inline.get("data").and_then(serde_json::Value::as_str) {
|
||||
return Ok(format!("data:{mime};base64,{data}"));
|
||||
}
|
||||
}
|
||||
Err("Gemini response missing inline image data".to_string())
|
||||
}
|
||||
|
||||
async fn generate_replicate(
|
||||
client: &reqwest::Client,
|
||||
prompt: &str,
|
||||
profile: &ImageGenProfile,
|
||||
width: Option<f64>,
|
||||
height: Option<f64>,
|
||||
) -> Result<String, String> {
|
||||
let base = profile
|
||||
.base_url
|
||||
.as_deref()
|
||||
.filter(|b| !b.trim().is_empty())
|
||||
.unwrap_or("https://api.replicate.com");
|
||||
let base = base.trim_end_matches('/');
|
||||
let mut input = serde_json::json!({ "prompt": prompt });
|
||||
if let Some(w) = width {
|
||||
input["width"] = serde_json::json!(w as i64);
|
||||
}
|
||||
if let Some(h) = height {
|
||||
input["height"] = serde_json::json!(h as i64);
|
||||
}
|
||||
let resp = client
|
||||
.post(format!("{base}/v1/predictions"))
|
||||
.bearer_auth(profile.api_key.trim())
|
||||
.json(&serde_json::json!({ "model": profile.model, "input": input }))
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("Replicate request failed: {e}"))?;
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
if !status.is_success() {
|
||||
return Err(provider_error("Replicate", status, &body));
|
||||
}
|
||||
let prediction: serde_json::Value =
|
||||
serde_json::from_str(&body).map_err(|e| format!("Replicate response parse: {e}"))?;
|
||||
let Some(id) = prediction.get("id").and_then(serde_json::Value::as_str) else {
|
||||
return Err("Replicate response missing prediction ID".to_string());
|
||||
};
|
||||
// Poll until terminal (TS: max 120 s, 2 s interval). The deadline is
|
||||
// wall-clock, not iteration-count: each poll also carries its own short
|
||||
// timeout so a black-holing endpoint can't stretch "60 polls" into
|
||||
// hours of a held connection thread (the shared client's 120 s
|
||||
// per-request timeout multiplies per iteration otherwise).
|
||||
let deadline = tokio::time::Instant::now() + Duration::from_secs(120);
|
||||
loop {
|
||||
tokio::time::sleep(Duration::from_secs(2)).await;
|
||||
// A poll never runs past the deadline: its own timeout is capped to
|
||||
// the time remaining, so the loop's total stays ~120 s instead of
|
||||
// "deadline + one full poll".
|
||||
let now = tokio::time::Instant::now();
|
||||
if now >= deadline {
|
||||
break;
|
||||
}
|
||||
let poll_timeout = Duration::from_secs(15).min(deadline - now);
|
||||
let resp = client
|
||||
.get(format!("{base}/v1/predictions/{id}"))
|
||||
.bearer_auth(profile.api_key.trim())
|
||||
.timeout(poll_timeout)
|
||||
.send()
|
||||
.await
|
||||
.map_err(|e| format!("Replicate poll request failed: {e}"))?;
|
||||
let status = resp.status();
|
||||
let body = resp.text().await.unwrap_or_default();
|
||||
if !status.is_success() {
|
||||
return Err(format!(
|
||||
"Replicate poll returned {}: {}",
|
||||
status.as_u16(),
|
||||
body.chars().take(ERROR_MESSAGE_CAP).collect::<String>()
|
||||
));
|
||||
}
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&body).map_err(|e| format!("Replicate poll parse: {e}"))?;
|
||||
match json.get("status").and_then(serde_json::Value::as_str) {
|
||||
Some("succeeded") => {
|
||||
let output = json.get("output");
|
||||
let url = output
|
||||
.and_then(|o| o.as_array())
|
||||
.and_then(|a| a.first())
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| output.and_then(serde_json::Value::as_str));
|
||||
return url
|
||||
.map(str::to_string)
|
||||
.ok_or_else(|| "Replicate succeeded but output is missing".to_string());
|
||||
}
|
||||
Some(s @ ("failed" | "canceled")) => {
|
||||
let detail = json
|
||||
.get("error")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("unknown error");
|
||||
return Err(format!("Replicate prediction {s}: {detail}"));
|
||||
}
|
||||
_ => {}
|
||||
}
|
||||
}
|
||||
Err("Replicate prediction timed out after 120 seconds".to_string())
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
fn state_with_profile(id: &str, api_key: &str) -> op_editor_core::EditorState {
|
||||
let mut state = op_editor_core::EditorState::default();
|
||||
state.editor_ui.agent_settings.image_gen_profiles = vec![ImageGenProfile {
|
||||
id: id.to_string(),
|
||||
name: "p".into(),
|
||||
provider: ImageGenProvider::OpenAi,
|
||||
api_key: api_key.to_string(),
|
||||
model: "gpt-image-1".into(),
|
||||
base_url: None,
|
||||
test_status: ImageTestStatus::Idle,
|
||||
}];
|
||||
state
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_generate_request_prefers_the_request_profile() {
|
||||
let state = state_with_profile("persisted", "sk-persisted");
|
||||
let req = parse_generate_request(
|
||||
r#"{"prompt":"a cat","width":1600,"height":900,
|
||||
"profile":{"provider":"gemini","model":"gemini-img","api_key":"sk-req"}}"#,
|
||||
&state,
|
||||
)
|
||||
.expect("parses");
|
||||
assert_eq!(req.prompt, "a cat");
|
||||
assert_eq!(req.width, Some(1600.0));
|
||||
assert_eq!(req.profile.provider, ImageGenProvider::Gemini);
|
||||
assert_eq!(req.profile.api_key, "sk-req");
|
||||
assert!(!req.custom_endpoint);
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_generate_request_falls_back_to_the_daemon_profile() {
|
||||
let state = state_with_profile("persisted", "sk-persisted");
|
||||
let req = parse_generate_request(r#"{"prompt":"a cat"}"#, &state).expect("parses");
|
||||
assert_eq!(req.profile.api_key, "sk-persisted");
|
||||
// No configured profile anywhere → explicit error.
|
||||
let empty = op_editor_core::EditorState::default();
|
||||
let err = parse_generate_request(r#"{"prompt":"a cat"}"#, &empty)
|
||||
.err()
|
||||
.expect("unconfigured daemon must error");
|
||||
assert_eq!(err, "image generation not configured");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_generate_request_flags_custom_endpoints_and_rejects_bad_input() {
|
||||
let state = op_editor_core::EditorState::default();
|
||||
let req = parse_generate_request(
|
||||
r#"{"prompt":"x","profile":{"provider":"custom","model":"m","api_key":"k",
|
||||
"base_url":"https://images.example.com"}}"#,
|
||||
&state,
|
||||
)
|
||||
.expect("parses");
|
||||
assert!(req.custom_endpoint);
|
||||
assert!(parse_generate_request(r#"{"prompt":""}"#, &state).is_err());
|
||||
assert!(parse_generate_request(
|
||||
r#"{"prompt":"x","profile":{"provider":"nope","model":"m","api_key":"k"}}"#,
|
||||
&state
|
||||
)
|
||||
.is_err());
|
||||
assert!(parse_generate_request(
|
||||
r#"{"prompt":"x","profile":{"provider":"openai","model":"m","api_key":""}}"#,
|
||||
&state
|
||||
)
|
||||
.is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn size_mappers_mirror_ts() {
|
||||
assert_eq!(openai_size(None, None), "1024x1024");
|
||||
assert_eq!(openai_size(Some(1600.0), Some(900.0)), "1792x1024");
|
||||
assert_eq!(openai_size(Some(900.0), Some(1600.0)), "1024x1792");
|
||||
assert_eq!(gemini_aspect_ratio(None, None), None);
|
||||
assert_eq!(
|
||||
gemini_aspect_ratio(Some(1920.0), Some(1080.0)),
|
||||
Some("16:9")
|
||||
);
|
||||
assert_eq!(gemini_aspect_ratio(Some(500.0), Some(500.0)), Some("1:1"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn provider_error_prefers_the_message_field() {
|
||||
let status = reqwest::StatusCode::BAD_GATEWAY;
|
||||
assert_eq!(
|
||||
provider_error(
|
||||
"OpenAI",
|
||||
status,
|
||||
r#"{"error":{"message":"quota exceeded"}}"#
|
||||
),
|
||||
"quota exceeded"
|
||||
);
|
||||
assert_eq!(
|
||||
provider_error("Replicate", status, r#"{"detail":"invalid token"}"#),
|
||||
"invalid token"
|
||||
);
|
||||
assert!(provider_error("Gemini", status, "<html>boom</html>")
|
||||
.starts_with("Gemini returned 502"));
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn error_json_truncates_to_the_ts_cap() {
|
||||
let long = "x".repeat(300);
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&generate_error_json(&long)).expect("valid json");
|
||||
assert_eq!(json["error"].as_str().unwrap().len(), ERROR_MESSAGE_CAP);
|
||||
let ok: serde_json::Value =
|
||||
serde_json::from_str(&generate_ok_json("data:image/png;base64,AA==")).unwrap();
|
||||
assert_eq!(ok["ok"], true);
|
||||
}
|
||||
}
|
||||
712
crates/op-host-services/src/web_image_search.rs
Normal file
712
crates/op-host-services/src/web_image_search.rs
Normal file
|
|
@ -0,0 +1,712 @@
|
|||
//! Browser image-search route for the web daemon.
|
||||
//!
|
||||
//! `POST /api/ai/image/search` mirrors the desktop image panel's Search
|
||||
//! popover backend (`op-host-desktop/src/image_panel_host.rs`: Openverse →
|
||||
//! two-keyword retry → Wikimedia, thumbnails embedded as `data:` URLs) so
|
||||
//! the wasm shell can drain its `search_epoch` through the daemon instead
|
||||
//! of leaving the popover loading forever. Openverse credentials come from
|
||||
//! the request body (browser-held) or fall back to the daemon's persisted
|
||||
//! agent settings. Openverse / Wikimedia are product-constant public hosts
|
||||
//! — the same operator-trust tier as the desktop path — so they dial with
|
||||
//! a plain client; nothing in this route dials a browser-supplied URL.
|
||||
//!
|
||||
//! Unlike the desktop, fetched thumbnails are NOT re-encoded/down-scaled
|
||||
//! here: `image_downscale` needs skia and this crate must stay GL-free for
|
||||
//! `op-host-web-server`. The 4 MiB per-image cap still bounds what can be
|
||||
//! embedded.
|
||||
|
||||
use std::time::Duration;
|
||||
|
||||
use reqwest::header::CONTENT_TYPE;
|
||||
|
||||
/// TS popover requests `count: 5` (desktop parity).
|
||||
const SEARCH_RESULT_COUNT: usize = 5;
|
||||
const MAX_EMBEDDED_IMAGE_BYTES: usize = 4 * 1024 * 1024;
|
||||
|
||||
/// Design-artifact words that are pure noise against a photo corpus (see
|
||||
/// the desktop `image_search_session.rs` for the measurement notes).
|
||||
const IMAGE_SEARCH_ARTIFACT_WORDS: &[&str] = &[
|
||||
"album",
|
||||
"cover",
|
||||
"playlist",
|
||||
"artwork",
|
||||
"poster",
|
||||
"thumbnail",
|
||||
"logo",
|
||||
"icon",
|
||||
"banner",
|
||||
"mockup",
|
||||
"screenshot",
|
||||
"wallpaper",
|
||||
];
|
||||
|
||||
const IMAGE_SEARCH_STOP_WORDS: &[&str] = &[
|
||||
"a",
|
||||
"an",
|
||||
"the",
|
||||
"and",
|
||||
"or",
|
||||
"but",
|
||||
"in",
|
||||
"on",
|
||||
"at",
|
||||
"to",
|
||||
"for",
|
||||
"of",
|
||||
"with",
|
||||
"by",
|
||||
"from",
|
||||
"is",
|
||||
"are",
|
||||
"was",
|
||||
"were",
|
||||
"be",
|
||||
"been",
|
||||
"being",
|
||||
"have",
|
||||
"has",
|
||||
"had",
|
||||
"do",
|
||||
"does",
|
||||
"did",
|
||||
"will",
|
||||
"would",
|
||||
"could",
|
||||
"should",
|
||||
"may",
|
||||
"might",
|
||||
"shall",
|
||||
"can",
|
||||
"that",
|
||||
"this",
|
||||
"these",
|
||||
"those",
|
||||
"it",
|
||||
"its",
|
||||
"very",
|
||||
"really",
|
||||
"just",
|
||||
"also",
|
||||
"about",
|
||||
"above",
|
||||
"after",
|
||||
"before",
|
||||
"between",
|
||||
"into",
|
||||
"through",
|
||||
"during",
|
||||
"each",
|
||||
"some",
|
||||
"such",
|
||||
"no",
|
||||
"not",
|
||||
"only",
|
||||
"same",
|
||||
"so",
|
||||
"than",
|
||||
"too",
|
||||
"up",
|
||||
"out",
|
||||
"if",
|
||||
"then",
|
||||
"once",
|
||||
"here",
|
||||
"there",
|
||||
"when",
|
||||
"where",
|
||||
"how",
|
||||
"all",
|
||||
"both",
|
||||
"few",
|
||||
"more",
|
||||
"most",
|
||||
"other",
|
||||
"any",
|
||||
"as",
|
||||
"while",
|
||||
"using",
|
||||
"showing",
|
||||
"featuring",
|
||||
"looking",
|
||||
"style",
|
||||
"styled",
|
||||
"inspired",
|
||||
"based",
|
||||
];
|
||||
|
||||
#[derive(Clone, PartialEq, Eq)]
|
||||
pub(crate) struct WebOpenverseCredentials {
|
||||
pub(crate) client_id: String,
|
||||
pub(crate) client_secret: String,
|
||||
}
|
||||
|
||||
impl WebOpenverseCredentials {
|
||||
fn from_parts(client_id: &str, client_secret: &str) -> Option<Self> {
|
||||
let client_id = client_id.trim();
|
||||
let client_secret = client_secret.trim();
|
||||
if client_id.is_empty() || client_secret.is_empty() {
|
||||
None
|
||||
} else {
|
||||
Some(Self {
|
||||
client_id: client_id.to_string(),
|
||||
client_secret: client_secret.to_string(),
|
||||
})
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
/// One search hit ready for the JSON reply.
|
||||
pub(crate) struct WebImageSearchHit {
|
||||
pub(crate) id: String,
|
||||
pub(crate) thumb_data_url: String,
|
||||
pub(crate) attribution: String,
|
||||
}
|
||||
|
||||
pub(crate) struct WebImageSearchOutcome {
|
||||
pub(crate) results: Vec<WebImageSearchHit>,
|
||||
/// `"openverse"` / `"wikimedia"`, `None` when nothing landed.
|
||||
pub(crate) source: Option<&'static str>,
|
||||
}
|
||||
|
||||
/// Parse the request body and snapshot the daemon-side credential fallback.
|
||||
/// Returns `(query, credentials)` or an error message for the 400 reply.
|
||||
pub(crate) fn parse_search_request(
|
||||
body: &str,
|
||||
state: &op_editor_core::EditorState,
|
||||
) -> Result<(String, Option<WebOpenverseCredentials>), String> {
|
||||
let value: serde_json::Value =
|
||||
serde_json::from_str(body).map_err(|_| "invalid request body".to_string())?;
|
||||
let obj = value.as_object().ok_or("invalid request body")?;
|
||||
let query = obj
|
||||
.get("query")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|q| !q.is_empty())
|
||||
.ok_or("missing query")?;
|
||||
// Browser-held credential wins; the daemon's persisted settings are the
|
||||
// fallback (both are optional — anonymous Openverse works, rate-limited).
|
||||
let request_credentials = obj
|
||||
.get("openverse")
|
||||
.and_then(serde_json::Value::as_object)
|
||||
.and_then(|cred| {
|
||||
WebOpenverseCredentials::from_parts(
|
||||
cred.get("client_id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(""),
|
||||
cred.get("client_secret")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(""),
|
||||
)
|
||||
});
|
||||
let credentials = request_credentials.or_else(|| {
|
||||
let settings = &state.editor_ui.agent_settings;
|
||||
WebOpenverseCredentials::from_parts(
|
||||
&settings.openverse_client_id,
|
||||
&settings.openverse_client_secret,
|
||||
)
|
||||
});
|
||||
Ok((query.to_string(), credentials))
|
||||
}
|
||||
|
||||
/// JSON reply body for a finished search.
|
||||
pub(crate) fn search_outcome_to_json(outcome: &WebImageSearchOutcome) -> String {
|
||||
let results: Vec<serde_json::Value> = outcome
|
||||
.results
|
||||
.iter()
|
||||
.map(|hit| {
|
||||
serde_json::json!({
|
||||
"id": hit.id,
|
||||
"thumb_data_url": hit.thumb_data_url,
|
||||
"attribution": hit.attribution,
|
||||
})
|
||||
})
|
||||
.collect();
|
||||
serde_json::json!({
|
||||
"ok": true,
|
||||
"results": results,
|
||||
"source": outcome.source,
|
||||
})
|
||||
.to_string()
|
||||
}
|
||||
|
||||
/// Run the full search ladder on the calling thread (the connection's own
|
||||
/// thread — the caller must NOT hold the state lock).
|
||||
pub(crate) fn run_search_blocking(
|
||||
query: &str,
|
||||
credentials: Option<&WebOpenverseCredentials>,
|
||||
) -> WebImageSearchOutcome {
|
||||
let empty = WebImageSearchOutcome {
|
||||
results: Vec::new(),
|
||||
source: None,
|
||||
};
|
||||
let Ok(runtime) = tokio::runtime::Builder::new_current_thread()
|
||||
.enable_all()
|
||||
.build()
|
||||
else {
|
||||
return empty;
|
||||
};
|
||||
runtime.block_on(run_search(query, credentials))
|
||||
}
|
||||
|
||||
async fn run_search(
|
||||
query: &str,
|
||||
credentials: Option<&WebOpenverseCredentials>,
|
||||
) -> WebImageSearchOutcome {
|
||||
let empty = WebImageSearchOutcome {
|
||||
results: Vec::new(),
|
||||
source: None,
|
||||
};
|
||||
let Ok(client) = reqwest::Client::builder()
|
||||
.timeout(Duration::from_secs(8))
|
||||
.user_agent(concat!("openpencil-web-daemon/", env!("CARGO_PKG_VERSION")))
|
||||
.build()
|
||||
else {
|
||||
return empty;
|
||||
};
|
||||
// Simplify verbose prompts into keywords (TS simplifySearchQuery).
|
||||
let query = simplify_search_query(query);
|
||||
|
||||
// Openverse first; a zero-result answer retries with the first two
|
||||
// keywords before falling through to Wikimedia (desktop parity).
|
||||
let mut hits = fetch_openverse_list(&client, &query, credentials).await;
|
||||
if hits.as_ref().is_some_and(Vec::is_empty) {
|
||||
if let Some(truncated) = two_keyword_retry(&query) {
|
||||
if let Some(retry) = fetch_openverse_list(&client, &truncated, credentials).await {
|
||||
if !retry.is_empty() {
|
||||
hits = Some(retry);
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
if let Some(urls) = hits.filter(|h| !h.is_empty()) {
|
||||
let results = materialize_thumbs(&client, urls).await;
|
||||
if !results.is_empty() {
|
||||
return WebImageSearchOutcome {
|
||||
results,
|
||||
source: Some("openverse"),
|
||||
};
|
||||
}
|
||||
}
|
||||
let mut wiki = fetch_wikimedia_list(&client, &query).await;
|
||||
if wiki.is_empty() {
|
||||
if let Some(truncated) = two_keyword_retry(&query) {
|
||||
wiki = fetch_wikimedia_list(&client, &truncated).await;
|
||||
}
|
||||
}
|
||||
let results = materialize_thumbs(&client, wiki).await;
|
||||
let source = (!results.is_empty()).then_some("wikimedia");
|
||||
WebImageSearchOutcome { results, source }
|
||||
}
|
||||
|
||||
fn two_keyword_retry(query: &str) -> Option<String> {
|
||||
let words: Vec<&str> = query.split_whitespace().filter(|w| !w.is_empty()).collect();
|
||||
(words.len() > 2).then(|| words[..2].join(" "))
|
||||
}
|
||||
|
||||
pub(crate) struct RawHit {
|
||||
id: String,
|
||||
thumb_url: String,
|
||||
attribution: String,
|
||||
}
|
||||
|
||||
/// `None` = request-level failure (429 / network), `Some([])` = the
|
||||
/// catalogue answered with zero hits (the ladder distinguishes the two).
|
||||
async fn fetch_openverse_list(
|
||||
client: &reqwest::Client,
|
||||
query: &str,
|
||||
credentials: Option<&WebOpenverseCredentials>,
|
||||
) -> Option<Vec<RawHit>> {
|
||||
let url = reqwest::Url::parse_with_params(
|
||||
"https://api.openverse.org/v1/images/",
|
||||
&[
|
||||
("q", query),
|
||||
("page_size", &SEARCH_RESULT_COUNT.to_string()),
|
||||
],
|
||||
)
|
||||
.ok()?;
|
||||
let mut request = client.get(url);
|
||||
if let Some(credentials) = credentials {
|
||||
if let Some(token) = fetch_openverse_token(client, credentials).await {
|
||||
request = request.bearer_auth(token);
|
||||
}
|
||||
}
|
||||
let resp = request.send().await.ok()?;
|
||||
if !resp.status().is_success() {
|
||||
return None;
|
||||
}
|
||||
let json: serde_json::Value = resp.json().await.ok()?;
|
||||
Some(parse_openverse_results(&json))
|
||||
}
|
||||
|
||||
pub(crate) fn parse_openverse_results(json: &serde_json::Value) -> Vec<RawHit> {
|
||||
let Some(results) = json.get("results").and_then(serde_json::Value::as_array) else {
|
||||
return Vec::new();
|
||||
};
|
||||
results
|
||||
.iter()
|
||||
.filter_map(|r| {
|
||||
let thumb = r
|
||||
.get("thumbnail")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| r.get("url").and_then(serde_json::Value::as_str))?;
|
||||
let license = format!(
|
||||
"{} {}",
|
||||
r.get("license")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(""),
|
||||
r.get("license_version")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or(""),
|
||||
);
|
||||
Some(RawHit {
|
||||
id: r
|
||||
.get("id")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or_default()
|
||||
.to_string(),
|
||||
thumb_url: thumb.to_string(),
|
||||
attribution: r
|
||||
.get("attribution")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::to_string)
|
||||
.unwrap_or_else(|| license.trim().to_string()),
|
||||
})
|
||||
})
|
||||
.take(SEARCH_RESULT_COUNT)
|
||||
.collect()
|
||||
}
|
||||
|
||||
async fn fetch_wikimedia_list(client: &reqwest::Client, query: &str) -> Vec<RawHit> {
|
||||
let Ok(url) = reqwest::Url::parse_with_params(
|
||||
"https://commons.wikimedia.org/w/api.php",
|
||||
&[
|
||||
("action", "query"),
|
||||
("generator", "search"),
|
||||
("gsrsearch", query),
|
||||
("gsrnamespace", "6"),
|
||||
("gsrlimit", &SEARCH_RESULT_COUNT.to_string()),
|
||||
("prop", "imageinfo"),
|
||||
("iiprop", "url|size|mime|extmetadata"),
|
||||
("iiurlwidth", "800"),
|
||||
("format", "json"),
|
||||
("origin", "*"),
|
||||
],
|
||||
) else {
|
||||
return Vec::new();
|
||||
};
|
||||
let Ok(resp) = client.get(url).send().await else {
|
||||
return Vec::new();
|
||||
};
|
||||
if !resp.status().is_success() {
|
||||
return Vec::new();
|
||||
}
|
||||
let Ok(json) = resp.json::<serde_json::Value>().await else {
|
||||
return Vec::new();
|
||||
};
|
||||
parse_wikimedia_results(&json)
|
||||
}
|
||||
|
||||
pub(crate) fn parse_wikimedia_results(json: &serde_json::Value) -> Vec<RawHit> {
|
||||
let Some(pages) = json
|
||||
.get("query")
|
||||
.and_then(|q| q.get("pages"))
|
||||
.and_then(serde_json::Value::as_object)
|
||||
else {
|
||||
return Vec::new();
|
||||
};
|
||||
pages
|
||||
.values()
|
||||
.filter_map(|page| {
|
||||
let info = page.get("imageinfo")?.as_array()?.first()?;
|
||||
let thumb = info
|
||||
.get("thumburl")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.or_else(|| info.get("url").and_then(serde_json::Value::as_str))?;
|
||||
Some(RawHit {
|
||||
id: page
|
||||
.get("pageid")
|
||||
.map(|v| v.to_string())
|
||||
.unwrap_or_default(),
|
||||
thumb_url: thumb.to_string(),
|
||||
attribution: info
|
||||
.get("extmetadata")
|
||||
.and_then(|m| m.get("LicenseShortName"))
|
||||
.and_then(|l| l.get("value"))
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.unwrap_or("")
|
||||
.to_string(),
|
||||
})
|
||||
})
|
||||
.take(SEARCH_RESULT_COUNT)
|
||||
.collect()
|
||||
}
|
||||
|
||||
/// Download each hit's thumbnail into a `data:` URL. Hits whose thumbnails
|
||||
/// fail to download are dropped.
|
||||
async fn materialize_thumbs(client: &reqwest::Client, hits: Vec<RawHit>) -> Vec<WebImageSearchHit> {
|
||||
let mut out = Vec::with_capacity(hits.len());
|
||||
for hit in hits {
|
||||
if let Some(data_url) = fetch_image_data_url(client, &hit.thumb_url).await {
|
||||
out.push(WebImageSearchHit {
|
||||
id: hit.id,
|
||||
thumb_data_url: data_url,
|
||||
attribution: hit.attribution,
|
||||
});
|
||||
}
|
||||
}
|
||||
out
|
||||
}
|
||||
|
||||
/// Simplify a verbose prompt into provider keywords (desktop parity).
|
||||
pub(crate) fn simplify_search_query(prompt: &str) -> String {
|
||||
let mut normalized = String::with_capacity(prompt.len());
|
||||
for ch in prompt.to_lowercase().chars() {
|
||||
if ch.is_ascii_alphanumeric() || ch.is_ascii_whitespace() || ch == '-' {
|
||||
normalized.push(ch);
|
||||
} else {
|
||||
normalized.push(' ');
|
||||
}
|
||||
}
|
||||
let keywords: Vec<&str> = normalized
|
||||
.split_whitespace()
|
||||
.filter(|word| word.len() > 2 && !IMAGE_SEARCH_STOP_WORDS.contains(word))
|
||||
.take(6)
|
||||
.collect();
|
||||
// Drop artifact words ONLY when aesthetic words remain — "logo" alone
|
||||
// must not become an empty query.
|
||||
let non_artifact: Vec<&str> = keywords
|
||||
.iter()
|
||||
.copied()
|
||||
.filter(|word| !IMAGE_SEARCH_ARTIFACT_WORDS.contains(word))
|
||||
.collect();
|
||||
let keywords: Vec<&str> = if non_artifact.is_empty() {
|
||||
keywords
|
||||
} else {
|
||||
non_artifact
|
||||
}
|
||||
.into_iter()
|
||||
.take(4)
|
||||
.collect();
|
||||
if keywords.is_empty() {
|
||||
prompt.chars().take(30).collect()
|
||||
} else {
|
||||
keywords.join(" ")
|
||||
}
|
||||
}
|
||||
|
||||
pub(crate) async fn fetch_openverse_token(
|
||||
client: &reqwest::Client,
|
||||
credentials: &WebOpenverseCredentials,
|
||||
) -> Option<String> {
|
||||
let resp = client
|
||||
.post("https://api.openverse.org/v1/auth_tokens/token/")
|
||||
.form(&[
|
||||
("grant_type", "client_credentials"),
|
||||
("client_id", credentials.client_id.as_str()),
|
||||
("client_secret", credentials.client_secret.as_str()),
|
||||
])
|
||||
.send()
|
||||
.await
|
||||
.ok()?;
|
||||
if !resp.status().is_success() {
|
||||
return None;
|
||||
}
|
||||
let json: serde_json::Value = resp.json().await.ok()?;
|
||||
json.get("access_token")
|
||||
.and_then(serde_json::Value::as_str)
|
||||
.map(str::trim)
|
||||
.filter(|token| !token.is_empty())
|
||||
.map(str::to_string)
|
||||
}
|
||||
|
||||
/// Download `url` and embed it as a `data:` URL, subject to the 4 MiB cap.
|
||||
pub(crate) async fn fetch_image_data_url(client: &reqwest::Client, url: &str) -> Option<String> {
|
||||
let resp = client.get(url).send().await.ok()?;
|
||||
if !resp.status().is_success() {
|
||||
return None;
|
||||
}
|
||||
if resp
|
||||
.content_length()
|
||||
.is_some_and(|len| len > MAX_EMBEDDED_IMAGE_BYTES as u64)
|
||||
{
|
||||
return None;
|
||||
}
|
||||
let header_mime = resp
|
||||
.headers()
|
||||
.get(CONTENT_TYPE)
|
||||
.and_then(|value| value.to_str().ok())
|
||||
.and_then(normalize_image_mime_header);
|
||||
let bytes = resp.bytes().await.ok()?;
|
||||
if bytes.is_empty() || bytes.len() > MAX_EMBEDDED_IMAGE_BYTES {
|
||||
return None;
|
||||
}
|
||||
let mime = header_mime.or_else(|| sniff_image_mime(&bytes).map(str::to_string))?;
|
||||
use base64::engine::general_purpose::STANDARD as B64;
|
||||
use base64::Engine as _;
|
||||
Some(format!("data:{mime};base64,{}", B64.encode(&bytes)))
|
||||
}
|
||||
|
||||
fn normalize_image_mime_header(value: &str) -> Option<String> {
|
||||
let mime = value.split(';').next()?.trim().to_ascii_lowercase();
|
||||
if mime == "image/jpg" {
|
||||
return Some("image/jpeg".to_string());
|
||||
}
|
||||
if mime.starts_with("image/") && mime != "image/svg+xml" {
|
||||
Some(mime)
|
||||
} else {
|
||||
None
|
||||
}
|
||||
}
|
||||
|
||||
fn sniff_image_mime(bytes: &[u8]) -> Option<&'static str> {
|
||||
if bytes.starts_with(b"\x89PNG\r\n\x1A\n") {
|
||||
return Some("image/png");
|
||||
}
|
||||
if bytes.starts_with(b"\xFF\xD8\xFF") {
|
||||
return Some("image/jpeg");
|
||||
}
|
||||
if bytes.starts_with(b"GIF87a") || bytes.starts_with(b"GIF89a") {
|
||||
return Some("image/gif");
|
||||
}
|
||||
if bytes.len() >= 12 && bytes.starts_with(b"RIFF") && &bytes[8..12] == b"WEBP" {
|
||||
return Some("image/webp");
|
||||
}
|
||||
None
|
||||
}
|
||||
|
||||
#[cfg(test)]
|
||||
mod tests {
|
||||
use super::*;
|
||||
|
||||
#[test]
|
||||
fn simplify_search_query_mirrors_the_desktop_adapter() {
|
||||
assert_eq!(
|
||||
simplify_search_query("A beautiful sunset over the mountains"),
|
||||
"beautiful sunset over mountains"
|
||||
);
|
||||
// Artifact words drop only when aesthetic words remain.
|
||||
assert_eq!(
|
||||
simplify_search_query("synthwave album cover neon"),
|
||||
"synthwave neon"
|
||||
);
|
||||
assert_eq!(simplify_search_query("logo"), "logo");
|
||||
// Empty keyword set falls back to a 30-char prefix.
|
||||
assert_eq!(simplify_search_query("の"), "の");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_search_request_reads_query_and_prefers_request_credentials() {
|
||||
let mut state = op_editor_core::EditorState::default();
|
||||
state.editor_ui.agent_settings.openverse_client_id = "persisted-id".into();
|
||||
state.editor_ui.agent_settings.openverse_client_secret = "persisted-secret".into();
|
||||
let (query, cred) = parse_search_request(
|
||||
r#"{"query":"cat","openverse":{"client_id":"req-id","client_secret":"req-secret"}}"#,
|
||||
&state,
|
||||
)
|
||||
.expect("parses");
|
||||
assert_eq!(query, "cat");
|
||||
assert_eq!(cred.expect("cred").client_id, "req-id");
|
||||
// No request credential → daemon-persisted fallback.
|
||||
let (_, cred) = parse_search_request(r#"{"query":"cat"}"#, &state).expect("parses");
|
||||
assert_eq!(cred.expect("cred").client_id, "persisted-id");
|
||||
// Neither → anonymous.
|
||||
let empty = op_editor_core::EditorState::default();
|
||||
let (_, cred) = parse_search_request(r#"{"query":"cat"}"#, &empty).expect("parses");
|
||||
assert!(cred.is_none());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_search_request_rejects_bad_bodies() {
|
||||
let state = op_editor_core::EditorState::default();
|
||||
assert!(parse_search_request("", &state).is_err());
|
||||
assert!(parse_search_request("{}", &state).is_err());
|
||||
assert!(parse_search_request(r#"{"query":" "}"#, &state).is_err());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_openverse_results_maps_thumbnail_license_and_cap() {
|
||||
let json = serde_json::json!({
|
||||
"results": [
|
||||
{"id": "a", "thumbnail": "https://x/a.jpg", "attribution": "By A"},
|
||||
{"id": "b", "url": "https://x/b.jpg", "license": "cc0", "license_version": "1.0"},
|
||||
{"id": "c"},
|
||||
{"id": "d", "thumbnail": "https://x/d.jpg"},
|
||||
{"id": "e", "thumbnail": "https://x/e.jpg"},
|
||||
{"id": "f", "thumbnail": "https://x/f.jpg"},
|
||||
{"id": "g", "thumbnail": "https://x/g.jpg"}
|
||||
]
|
||||
});
|
||||
let hits = parse_openverse_results(&json);
|
||||
assert_eq!(hits.len(), SEARCH_RESULT_COUNT); // "c" dropped, capped at 5
|
||||
assert_eq!(hits[0].id, "a");
|
||||
assert_eq!(hits[0].attribution, "By A");
|
||||
assert_eq!(hits[1].thumb_url, "https://x/b.jpg");
|
||||
assert_eq!(hits[1].attribution, "cc0 1.0");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn parse_wikimedia_results_maps_thumburl_and_license() {
|
||||
let json = serde_json::json!({
|
||||
"query": {"pages": {
|
||||
"1": {"pageid": 1, "imageinfo": [{
|
||||
"thumburl": "https://c/w1.jpg",
|
||||
"extmetadata": {"LicenseShortName": {"value": "CC BY-SA 4.0"}}
|
||||
}]},
|
||||
"2": {"pageid": 2, "imageinfo": [{"url": "https://c/w2.jpg"}]},
|
||||
"3": {"pageid": 3}
|
||||
}}
|
||||
});
|
||||
let mut hits = parse_wikimedia_results(&json);
|
||||
hits.sort_by(|a, b| a.id.cmp(&b.id));
|
||||
assert_eq!(hits.len(), 2);
|
||||
assert_eq!(hits[0].thumb_url, "https://c/w1.jpg");
|
||||
assert_eq!(hits[0].attribution, "CC BY-SA 4.0");
|
||||
assert_eq!(hits[1].thumb_url, "https://c/w2.jpg");
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn search_outcome_json_shape() {
|
||||
let outcome = WebImageSearchOutcome {
|
||||
results: vec![WebImageSearchHit {
|
||||
id: "a".into(),
|
||||
thumb_data_url: "data:image/png;base64,AA==".into(),
|
||||
attribution: "By A".into(),
|
||||
}],
|
||||
source: Some("openverse"),
|
||||
};
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&search_outcome_to_json(&outcome)).expect("valid json");
|
||||
assert_eq!(json["ok"], true);
|
||||
assert_eq!(json["source"], "openverse");
|
||||
assert_eq!(json["results"][0]["id"], "a");
|
||||
assert_eq!(
|
||||
json["results"][0]["thumb_data_url"],
|
||||
"data:image/png;base64,AA=="
|
||||
);
|
||||
let empty = WebImageSearchOutcome {
|
||||
results: Vec::new(),
|
||||
source: None,
|
||||
};
|
||||
let json: serde_json::Value =
|
||||
serde_json::from_str(&search_outcome_to_json(&empty)).expect("valid json");
|
||||
assert!(json["source"].is_null());
|
||||
}
|
||||
|
||||
#[test]
|
||||
fn sniff_image_mime_recognizes_the_embeddable_formats() {
|
||||
assert_eq!(sniff_image_mime(b"\x89PNG\r\n\x1A\nxx"), Some("image/png"));
|
||||
assert_eq!(sniff_image_mime(b"\xFF\xD8\xFFxx"), Some("image/jpeg"));
|
||||
assert_eq!(sniff_image_mime(b"GIF89a"), Some("image/gif"));
|
||||
assert_eq!(
|
||||
sniff_image_mime(b"RIFF\0\0\0\0WEBPVP8 "),
|
||||
Some("image/webp")
|
||||
);
|
||||
assert_eq!(sniff_image_mime(b"<svg>"), None);
|
||||
assert_eq!(
|
||||
normalize_image_mime_header("image/jpg"),
|
||||
Some("image/jpeg".into())
|
||||
);
|
||||
assert_eq!(normalize_image_mime_header("image/svg+xml"), None);
|
||||
assert_eq!(normalize_image_mime_header("text/html"), None);
|
||||
}
|
||||
}
|
||||
Loading…
Reference in a new issue