1use crate::api::core_capabilities::handle_core_capability;
2use crate::api::openai_compat::AppState;
3use crate::api::package_ws::dispatch_package_payload;
4use axum::extract::{Path, State};
5use axum::http::StatusCode;
6use axum::Json;
7use std::path::{Path as FsPath, PathBuf};
8
9async fn download_image_to_workspace(url: &str) -> Option<String> {
13 let client = match reqwest::Client::builder()
17 .use_native_tls()
18 .timeout(std::time::Duration::from_secs(60))
19 .build()
20 {
21 Ok(c) => c,
22 Err(e) => {
23 tracing::warn!("image download client build failed: {e}");
24 return None;
25 }
26 };
27 let resp = match client.get(url).send().await {
28 Ok(r) => r,
29 Err(e) => {
30 tracing::warn!("image download failed (send): {e}");
31 return None;
32 }
33 };
34 let bytes = match resp.bytes().await {
35 Ok(b) => b,
36 Err(e) => {
37 tracing::warn!("image download failed (bytes): {e}");
38 return None;
39 }
40 };
41 tracing::info!("image downloaded: {} bytes from {}", bytes.len(), &url[..url.len().min(60)]);
42 let ext = url
44 .split('?')
45 .next()
46 .and_then(|u| u.rsplit('.').next())
47 .filter(|e| matches!(e.to_lowercase().as_str(), "png" | "jpg" | "jpeg" | "webp" | "gif"))
48 .unwrap_or("png")
49 .to_string();
50 let dir = std::path::Path::new("./workspace/image-gen");
51 let _ = std::fs::create_dir_all(dir);
52 let stamp = std::time::SystemTime::now()
54 .duration_since(std::time::UNIX_EPOCH)
55 .map(|d| d.as_nanos())
56 .unwrap_or(0);
57 let fname = format!("dl-{}-{}.{}", bytes.len(), stamp, ext);
58 let full = dir.join(&fname);
59 std::fs::write(&full, &bytes).ok()?;
60 let abs = std::fs::canonicalize(&full)
61 .map(|p| {
62 let s = p.to_string_lossy().to_string();
63 s.strip_prefix(r"\\?\").map(|x| x.to_string()).unwrap_or(s)
64 })
65 .unwrap_or_else(|_| full.to_string_lossy().to_string());
66 Some(abs)
67}
68
69async fn workspace_root_for_app(state: &AppState, app_name: Option<&str>) -> Option<PathBuf> {
70 let app_name = app_name?;
71 let apps = state.resolved_apps.read().await;
72 let app = apps.get(app_name)?;
73 if !app.sources.manifest_path.is_empty() {
74 let manifest_dir = FsPath::new(&app.sources.manifest_path).parent()?;
75 let candidate = manifest_dir.join("workspace");
76 if candidate.exists() {
77 return Some(candidate);
78 }
79 }
80 let config_path = app.config_path.as_ref()?;
81 let app_dir = FsPath::new(config_path).parent()?;
82 let config = crate::app::load_app_config(app_dir).ok()?;
83 let workspace = config.app_runtime.workspace;
84 if workspace.is_empty() {
85 return None;
86 }
87
88 Some(if FsPath::new(&workspace).is_absolute() {
89 PathBuf::from(workspace)
90 } else {
91 app_dir.join(workspace)
92 })
93}
94
95fn package_for_provider<'a>(
96 state: &'a AppState,
97 provider: &str,
98) -> Option<&'a crate::app::PackageSource> {
99 state.package_index.get(provider)
100}
101
102async fn enforce_provider_security(
103 state: &AppState,
104 provider: &str,
105 provider_runtime: &str,
106) -> Result<(), (StatusCode, serde_json::Value)> {
107 let profile = *state.active_profile.read().await;
108 if provider_runtime == "core" {
109 return Ok(());
110 }
111
112 let Some(pkg) = package_for_provider(state, provider) else {
113 return Err((
114 StatusCode::FORBIDDEN,
115 serde_json::json!({
116 "error": format!("Provider '{}' is not present in package index", provider),
117 "provider": provider,
118 "reason": "provider_missing_from_index",
119 }),
120 ));
121 };
122
123 let signature_ok = match profile {
124 crate::app::AppProfile::Safe => {
125 pkg.signature.starts_with("builtin:")
126 || (pkg.signature.starts_with("ed25519:")
127 && crate::app::verify_package_signature_for_source(
128 &pkg.signature,
129 &crate::app::signature_message(
130 &pkg.name,
131 "current",
132 &crate::api::generations::package_digest(
133 &state.repo_root,
134 &pkg.current_source,
135 ),
136 &pkg.current_source,
137 ),
138 &pkg.source_authority,
139 &pkg.source_public_keys,
140 )
141 .is_ok())
142 }
143 crate::app::AppProfile::Developer => {
144 (!pkg.signature.is_empty() && pkg.signature != "unsigned")
145 || (pkg.signature.starts_with("ed25519:")
146 && crate::app::verify_package_signature_for_source(
147 &pkg.signature,
148 &crate::app::signature_message(
149 &pkg.name,
150 "current",
151 &crate::api::generations::package_digest(
152 &state.repo_root,
153 &pkg.current_source,
154 ),
155 &pkg.current_source,
156 ),
157 &pkg.source_authority,
158 &pkg.source_public_keys,
159 )
160 .is_ok())
161 }
162 crate::app::AppProfile::Trusted => true,
163 };
164 if !pkg.trusted && profile == crate::app::AppProfile::Safe {
165 return Err((
166 StatusCode::FORBIDDEN,
167 serde_json::json!({
168 "error": format!("Provider '{}' is not trusted under safe profile", provider),
169 "provider": provider,
170 "signature": pkg.signature,
171 "reason": "provider_not_trusted",
172 }),
173 ));
174 }
175 if !signature_ok {
176 return Err((
177 StatusCode::FORBIDDEN,
178 serde_json::json!({
179 "error": format!("Provider '{}' signature '{}' is not accepted under profile '{}'", provider, pkg.signature, profile.as_str()),
180 "provider": provider,
181 "signature": pkg.signature,
182 "profile": profile.as_str(),
183 "reason": "signature_rejected",
184 }),
185 ));
186 }
187
188 Ok(())
189}
190
191pub async fn execute_capability_call(
192 state: &AppState,
193 name: &str,
194 payload: serde_json::Value,
195) -> Result<serde_json::Value, (StatusCode, serde_json::Value)> {
196 let registry = state.capability_registry.read().await;
197 let capability = if let Some(capability) = registry.get(name) {
198 capability.clone()
199 } else {
200 return Err((
201 StatusCode::NOT_FOUND,
202 serde_json::json!({
203 "error": format!("Capability '{}' not found", name)
204 }),
205 ));
206 };
207 drop(registry);
208
209 {
210 let profile = *state.active_profile.read().await;
211 let decision = state.core_policy.check(name, profile);
212 if !decision.allowed {
213 return Err((
214 StatusCode::FORBIDDEN,
215 serde_json::json!({
216 "error": format!("Policy denied: {}", decision.reason),
217 "capability": name,
218 "profile": profile.as_str(),
219 }),
220 ));
221 }
222 }
223
224 let selected_provider = payload
225 .get("provider")
226 .and_then(|value| value.as_str())
227 .map(|value| value.to_string())
228 .or_else(|| {
229 let app_name = payload.get("app").and_then(|value| value.as_str())?;
230 capability
231 .bindings
232 .iter()
233 .find(|binding| binding.app == app_name)
234 .map(|binding| binding.provider.clone())
235 })
236 .or_else(|| {
237 capability
238 .providers
239 .first()
240 .map(|provider| provider.provider.clone())
241 });
242
243 let provider = if let Some(p) = selected_provider {
244 p
245 } else {
246 return Err((
247 StatusCode::BAD_REQUEST,
248 serde_json::json!({
249 "error": format!("Capability '{}' has no available provider", name)
250 }),
251 ));
252 };
253
254 let action = payload
255 .get("action")
256 .and_then(|v| v.as_str())
257 .unwrap_or("call");
258 let mut data = payload
259 .get("data")
260 .cloned()
261 .unwrap_or_else(|| serde_json::Value::Object(serde_json::Map::new()));
262
263 if name == "image.generate" {
268 if let serde_json::Value::Object(map) = &mut data {
269 let has_key = map
270 .get("api_key")
271 .and_then(|v| v.as_str())
272 .map(|s| !s.trim().is_empty())
273 .unwrap_or(false);
274 if !has_key {
275 let config = state.config.read().await;
276 let img_provider = config
277 .routing
278 .image_provider
279 .as_deref()
280 .and_then(|name| config.providers.iter().find(|p| p.name == name))
281 .or_else(|| {
282 config
283 .providers
284 .iter()
285 .find(|p| p.name.to_lowercase().contains("image"))
286 })
287 .or_else(|| {
290 config
291 .providers
292 .iter()
293 .find(|p| p.keys.iter().any(|k| k.enabled && !k.value.trim().is_empty()))
294 });
295 if let Some(p) = img_provider {
296 if let Some(k) = p.keys.iter().find(|k| k.enabled && !k.value.trim().is_empty()) {
297 map.insert(
298 "api_key".to_string(),
299 serde_json::Value::String(k.value.trim().to_string()),
300 );
301 }
302 if !p.base_url.trim().is_empty()
303 && map
304 .get("base_url")
305 .and_then(|v| v.as_str())
306 .map(|s| s.trim().is_empty())
307 .unwrap_or(true)
308 {
309 map.insert(
310 "base_url".to_string(),
311 serde_json::Value::String(p.base_url.trim().to_string()),
312 );
313 }
314 }
315 }
316 }
317 }
318
319 let provider_runtime = capability
320 .providers
321 .iter()
322 .find(|p| p.provider == provider)
323 .map(|p| p.runtime.as_str())
324 .unwrap_or("wasm");
325
326 enforce_provider_security(state, &provider, provider_runtime).await?;
327
328 if provider_runtime == "core" {
329 let app_name = payload.get("app").and_then(|value| value.as_str());
330 let workspace_root = workspace_root_for_app(state, app_name).await;
331 match handle_core_capability(name, action, &data, workspace_root.as_deref()).await {
332 Ok(response) => Ok(serde_json::json!({
333 "capability": name,
334 "provider": provider,
335 "response": response,
336 "status": "executed",
337 "mode": "core"
338 })),
339 Err(err) => Err((
340 StatusCode::INTERNAL_SERVER_ERROR,
341 serde_json::json!({
342 "error": err,
343 "capability": name,
344 "provider": provider,
345 }),
346 )),
347 }
348 } else if provider_runtime == "native" {
349 let native_handle = state.native_handle.read().await;
350 let Some(handle) = native_handle.as_ref() else {
351 return Err((
352 StatusCode::NOT_IMPLEMENTED,
353 serde_json::json!({
354 "error": format!(
355 "Native provider '{}' for capability '{}' is recognized but no native runtime is active.",
356 provider, name
357 ),
358 "capability": name,
359 "provider": provider,
360 "mode": "native-stub",
361 }),
362 ));
363 };
364
365 let native_payload = serde_json::json!({
366 "capability": name,
367 "action": action,
368 "data": data,
369 "app": payload.get("app").cloned().unwrap_or(serde_json::Value::Null),
370 });
371
372 match handle.call_json(&provider, &native_payload) {
373 Ok(response) => Ok(serde_json::json!({
374 "capability": name,
375 "provider": provider,
376 "response": response,
377 "status": "executed",
378 "mode": "native"
379 })),
380 Err(err) => Err((
381 StatusCode::BAD_GATEWAY,
382 serde_json::json!({
383 "error": err.to_string(),
384 "capability": name,
385 "provider": provider,
386 "mode": "native",
387 }),
388 )),
389 }
390 } else {
391 let envelope = serde_json::json!({
392 "action": action,
393 "data": data,
394 });
395
396 let mut response = dispatch_package_payload(&provider, envelope, state).await;
397
398 if name == "image.generate" {
401 let maybe_url = response
402 .get("data")
403 .and_then(|d| d.get("url"))
404 .and_then(|u| u.as_str())
405 .map(|s| s.to_string());
406 tracing::info!("image.generate postprocess: url present = {}", maybe_url.is_some());
407 if let Some(url) = maybe_url {
408 if let Some(local_path) = download_image_to_workspace(&url).await {
409 if let Some(data_obj) = response.get_mut("data").and_then(|d| d.as_object_mut()) {
410 data_obj.remove("url");
411 data_obj.insert("output_path".to_string(), serde_json::Value::String(local_path));
412 }
413 }
414 }
415 }
416
417 if response.get("error").is_some() {
418 Err((StatusCode::BAD_GATEWAY, response))
419 } else {
420 Ok(serde_json::json!({
421 "capability": name,
422 "provider": provider,
423 "response": response,
424 "status": "executed",
425 "mode": if provider_runtime == "service" { "service" } else { "wasm-phase" }
426 }))
427 }
428 }
429}
430
431pub async fn list_capabilities(State(state): State<AppState>) -> Json<serde_json::Value> {
432 let registry = state.capability_registry.read().await;
433 let values: Vec<_> = registry.values().cloned().collect();
434 Json(serde_json::json!({ "capabilities": values }))
435}
436
437pub async fn get_capability(
438 Path(name): Path<String>,
439 State(state): State<AppState>,
440) -> Result<Json<serde_json::Value>, (StatusCode, Json<serde_json::Value>)> {
441 let registry = state.capability_registry.read().await;
442 if let Some(capability) = registry.get(&name) {
443 Ok(Json(serde_json::json!({ "capability": capability })))
444 } else {
445 Err((
446 StatusCode::NOT_FOUND,
447 Json(serde_json::json!({
448 "error": format!("Capability '{}' not found", name)
449 })),
450 ))
451 }
452}
453
454pub async fn capability_call(
455 Path(name): Path<String>,
456 State(state): State<AppState>,
457 Json(payload): Json<serde_json::Value>,
458) -> Result<Json<serde_json::Value>, (StatusCode, Json<serde_json::Value>)> {
459 execute_capability_call(&state, &name, payload)
460 .await
461 .map(Json)
462 .map_err(|(status, value)| (status, Json(value)))
463}
464
465#[cfg(test)]
466mod tests {
467 use super::execute_capability_call;
468 use crate::api::openai_compat::AppState;
469 use crate::app::{
470 sign_package_message, signature_message, AppProfile, CapabilityProviderRecord,
471 CapabilityRegistry, CapabilityRegistryEntry, CorePolicy, GenerationStoreMap, PackageIndex,
472 PackageSource,
473 };
474 use crate::config::{
475 AppConfig, CoreConfig, FallbackConfig, KeyStrategyConfig, RegistryConfig, RoutingConfig,
476 };
477 use crate::defaults::{
478 error_handler::DefaultErrorHandler, key_selectors::FailoverSelector, router::DefaultRouter,
479 };
480 use crate::pipeline::Pipeline;
481 use crate::process::ProcessManager;
482 use crate::vkeys::VirtualKeyStore;
483 use axum::http::StatusCode;
484 use ed25519_dalek::SigningKey;
485 use std::collections::HashMap;
486 use std::sync::{Arc, Mutex as StdMutex};
487 use tokio::sync::RwLock;
488
489 fn test_state(
490 repo_root: std::path::PathBuf,
491 capability_registry: CapabilityRegistry,
492 package_index: PackageIndex,
493 profile: AppProfile,
494 ) -> AppState {
495 AppState {
496 config: Arc::new(RwLock::new(AppConfig {
497 core: CoreConfig::default(),
498 providers: vec![],
499 routing: RoutingConfig::default(),
500 key_strategy: KeyStrategyConfig::default(),
501 fallback: FallbackConfig::default(),
502 virtual_keys: vec![],
503 services: vec![],
504 packages: vec![],
505 registry: RegistryConfig::default(),
506 package_aliases: HashMap::new(),
507 web_search: Default::default(),
508 team: Default::default(),
509 })),
510 config_path: repo_root.join("config").join("config.toml"),
511 pipeline: Arc::new(Pipeline {
512 router: Arc::new(DefaultRouter {
513 default_provider: "".into(),
514 }),
515 key_selector: Arc::new(FailoverSelector),
516 transforms: Arc::new(crate::defaults::transforms::TransformRegistry::with_defaults()),
517 error_handler: Arc::new(DefaultErrorHandler { max_retries: 0 }),
518 http_client: reqwest::Client::new(),
519 }),
520 process_manager: Arc::new(ProcessManager::new()),
521 vkey_store: Arc::new(VirtualKeyStore::new()),
522 package_manager: Arc::new(RwLock::new(crate::package::PackageManager::new())),
523 wasm_handle: Arc::new(RwLock::new(None)),
524 native_handle: Arc::new(RwLock::new(None)),
525 resolved_apps: Arc::new(RwLock::new(Default::default())),
526 capability_registry: Arc::new(RwLock::new(capability_registry)),
527 active_profile: Arc::new(RwLock::new(profile)),
528 core_policy: Arc::new(CorePolicy::default_policy()),
529 generation_store: Arc::new(RwLock::new(GenerationStoreMap::new())),
530 package_index: Arc::new(package_index),
531 data_dir: repo_root.join("data"),
532 repo_root,
533 runtime_token: None,
534 runtime_token_path: None,
535 chat_providers: Arc::new(RwLock::new(vec![])),
536 shutdown_tx: Arc::new(StdMutex::new(None)),
537 stream_buffer: Arc::new(StdMutex::new(std::collections::HashMap::new())),
538 }
539 }
540
541 fn capability_registry_for(provider: &str, runtime: &str) -> CapabilityRegistry {
542 let mut registry = CapabilityRegistry::new();
543 registry.insert(
544 "cap.test".into(),
545 CapabilityRegistryEntry {
546 capability: "cap.test".into(),
547 providers: vec![CapabilityProviderRecord {
548 provider: provider.into(),
549 runtime: runtime.into(),
550 priority: 0,
551 }],
552 bindings: vec![],
553 },
554 );
555 registry
556 }
557
558 fn signed_package_source(
559 repo_root: &std::path::Path,
560 name: &str,
561 current_source: &str,
562 trusted: bool,
563 signature_digest: &str,
564 ) -> PackageSource {
565 let signing_key = SigningKey::from_bytes(&[9; 32]);
566 let message = signature_message(name, "current", signature_digest, current_source);
567 let signature = sign_package_message(&signing_key, &message);
568 let source_public_key = signature
569 .split(':')
570 .nth(1)
571 .expect("public key segment exists")
572 .to_string();
573
574 let source_dir = repo_root.join(current_source);
575 std::fs::create_dir_all(&source_dir).expect("source directory created");
576 std::fs::write(source_dir.join("package.toml"), "name = 'cap-provider'\n")
577 .expect("marker file written");
578
579 PackageSource {
580 name: name.into(),
581 kind: "wasm".into(),
582 package_kind: String::new(),
583 runtime_provider: name.into(),
584 current_source: current_source.into(),
585 trusted,
586 signature,
587 source_authority: "test-authority".into(),
588 source_public_keys: vec![source_public_key],
589 provides: vec![],
590 requires: vec![],
591 }
592 }
593
594 #[tokio::test]
595 async fn execute_capability_call_rejects_provider_with_mismatched_digest_signature() {
596 let repo_root = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
597 .join("..")
598 .join("target")
599 .join("test-security-mismatched-digest");
600 let provider = "cap-provider";
601 let source = "fixtures/cap-provider";
602 let package_index = PackageIndex {
603 version: 1,
604 revision: "test-rev".into(),
605 source_url: "local://packages".into(),
606 package_sources: vec![signed_package_source(
607 &repo_root, provider, source, true, "deadbeef",
608 )],
609 };
610 let state = test_state(
611 repo_root.clone(),
612 capability_registry_for(provider, "service"),
613 package_index,
614 AppProfile::Safe,
615 );
616
617 let error = execute_capability_call(
618 &state,
619 "cap.test",
620 serde_json::json!({
621 "action": "health",
622 "data": {},
623 "provider": provider,
624 }),
625 )
626 .await
627 .expect_err(
628 "provider should be rejected when signed digest does not match current source digest",
629 );
630
631 assert_eq!(error.0, StatusCode::FORBIDDEN);
632 assert!(error.1["error"]
633 .as_str()
634 .expect("error string")
635 .contains("is not accepted under profile 'safe'"));
636 }
637
638 #[tokio::test]
639 async fn execute_capability_call_rejects_untrusted_provider_under_safe_profile() {
640 let repo_root = std::path::PathBuf::from(env!("CARGO_MANIFEST_DIR"))
641 .join("..")
642 .join("target")
643 .join("test-security-untrusted-provider");
644 let provider = "cap-provider";
645 let source = "fixtures/cap-provider";
646 let source_dir = repo_root.join(source);
647 std::fs::create_dir_all(&source_dir).expect("source directory created");
648 std::fs::write(source_dir.join("package.toml"), "name = 'cap-provider'\n")
649 .expect("marker file written");
650 let digest = crate::api::generations::package_digest(&repo_root, source);
651 let package_index = PackageIndex {
652 version: 1,
653 revision: "test-rev".into(),
654 source_url: "local://packages".into(),
655 package_sources: vec![signed_package_source(
656 &repo_root, provider, source, false, &digest,
657 )],
658 };
659 let state = test_state(
660 repo_root,
661 capability_registry_for(provider, "service"),
662 package_index,
663 AppProfile::Safe,
664 );
665
666 let error = execute_capability_call(
667 &state,
668 "cap.test",
669 serde_json::json!({
670 "action": "health",
671 "data": {},
672 "provider": provider,
673 }),
674 )
675 .await
676 .expect_err("untrusted provider should be rejected under safe profile");
677
678 assert_eq!(error.0, StatusCode::FORBIDDEN);
679 assert!(error.1["error"]
680 .as_str()
681 .expect("error string")
682 .contains("is not trusted under safe profile"));
683 }
684}