Skip to main content

weft_core/api/
capabilities.rs

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
9/// 下载远程图片到 ./workspace/image-gen/ 并返回绝对路径(剥掉 \\?\ 前缀,
10/// 与 image-gen 本地保存格式一致,前端 mediaUrl 可解析为 /media URL)。
11/// 失败返回 None(调用方回退到原始 url)。
12async fn download_image_to_workspace(url: &str) -> Option<String> {
13    // 用 native-tls(Windows schannel)而非默认 rustls:rustls 的 HTTP/2 帧解析
14    // 与部分 CDN(storage.fonedis.cc) 不兼容,下载大二进制报 "error decoding response body"。
15    // schannel 与 curl 行为一致,可正常下载。
16    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    // 从 url 推断扩展名,默认 png
43    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    // 用内容哈希 + 时间避免重名(不用 rng,用 bytes 长度 + 纳秒)
53    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    // 图像生成:若调用方未显式提供 api_key/base_url,则从配置的图像 provider
264    // (routing.image_provider 指定;缺省回退第一个名字含 "image" 的 provider)
265    // 取 base_url + 首个 key 注入 data。使 image-gen WASM 无需依赖环境变量,
266    // 且改 provider 后下次调用即生效(无需重启 Core)。key 不经前端,留在 Core 内。
267    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                    // 最终回退:用第一个有可用 key 的 provider(不限制必须是"图像"provider,
288                    // 让画布生成不依赖 routing.image_provider 配置)。
289                    .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        // 图像生成后处理:若返回远程 url(部分模型如 flux 返回直链而非本地路径),
399        // 下载保存到 workspace 并替换为 output_path,统一走本地 /media 加载(更快+可缓存)。
400        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}