diff --git a/.gitignore b/.gitignore new file mode 100644 index 0000000..0253060 --- /dev/null +++ b/.gitignore @@ -0,0 +1,9 @@ +target/ +node_modules/ +**/*.log +.env +.env.local +coverage/ +.DS_Store +.idea/ +.vscode/ diff --git a/Cargo.toml b/Cargo.toml new file mode 100644 index 0000000..c96da19 --- /dev/null +++ b/Cargo.toml @@ -0,0 +1,45 @@ +[workspace] +members = [ + "apps/api", + "apps/auth", + "sandbox", + "tests" +] +resolver = "2" + +[workspace.package] +edition = "2021" + +[workspace.dependencies] +anyhow = "1.0" +async-trait = "0.1" +axum = { version = "0.7", features = ["macros", "ws"] } +bcrypt = "0.15" +bytes = "1.6" +chrono = { version = "0.4", features = ["serde"] } +hex = "0.4" +jsonwebtoken = "9.2" +opentelemetry = { version = "0.21", features = ["rt-tokio"] } +opentelemetry-otlp = { version = "0.14", features = ["grpc-tonic"] } +opentelemetry-prometheus = "0.13" +parking_lot = "0.12" +prometheus = "0.13" +rand = "0.8" +reqwest = { version = "0.12", default-features = false, features = ["json", "gzip", "rustls-tls"] } +serde = { version = "1.0", features = ["derive"] } +serde_json = "1.0" +serde_with = "3.8" +sha2 = "0.10" +sqlx = { version = "0.7", default-features = false, features = ["runtime-tokio", "postgres", "macros", "uuid", "chrono", "json"] } +thiserror = "1.0" +tokio = { version = "1.37", features = ["rt-multi-thread", "macros", "time", "process", "io-util", "fs"] } +tokio-util = { version = "0.7", features = ["sync"] } +tower = "0.4" +tower-http = { version = "0.5", features = ["trace", "cors"] } +tracing = "0.1" +tracing-opentelemetry = "0.22" +tracing-subscriber = { version = "0.3", features = ["env-filter", "json"] } +uuid = { version = "1.7", features = ["v4", "serde"] } +validator = "0.16" +tempfile = "3.10" +wat = "1.0" diff --git a/apps/api/Cargo.toml b/apps/api/Cargo.toml new file mode 100644 index 0000000..7370b7c --- /dev/null +++ b/apps/api/Cargo.toml @@ -0,0 +1,30 @@ +[package] +name = "api" +version = "0.1.0" +edition = "2021" + +[dependencies] +anyhow = { workspace = true } +axum = { workspace = true } +base64 = "0.22" +chrono = { workspace = true } +hex = { workspace = true } +jsonwebtoken = { workspace = true } +parking_lot = { workspace = true } +sha2 = { workspace = true } +reqwest = { workspace = true } +sqlx = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } +tokio = { workspace = true } +tracing = { workspace = true } +tracing-opentelemetry = { workspace = true } +tracing-subscriber = { workspace = true } +tower = { workspace = true } +tower-http = { workspace = true } +sandbox = { path = "../../sandbox" } +uuid = { workspace = true } +opentelemetry = { workspace = true } +opentelemetry-otlp = { workspace = true } +opentelemetry-prometheus = { workspace = true } +prometheus = { workspace = true } diff --git a/apps/api/src/main.rs b/apps/api/src/main.rs new file mode 100644 index 0000000..dfe1bbd --- /dev/null +++ b/apps/api/src/main.rs @@ -0,0 +1,2641 @@ +use std::net::SocketAddr; +use std::path::{Component, Path, PathBuf}; +use std::sync::Arc; +use std::time::{Duration, Instant}; + +use axum::extract::State; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; +use axum::routing::{get, post}; +use axum::{Json, Router}; +use base64::engine::general_purpose::STANDARD as BASE64; +use base64::Engine; +use chrono::{DateTime, Utc}; +use hex::encode as hex_encode; +use jsonwebtoken::{decode, Algorithm, DecodingKey, Validation}; +use opentelemetry::global; +use opentelemetry::metrics::{Counter, Histogram, UpDownCounter}; +use opentelemetry::sdk::trace::{self, Config}; +use opentelemetry::sdk::Resource; +use opentelemetry::KeyValue; +use opentelemetry_otlp::WithExportConfig; +use opentelemetry_prometheus::PrometheusExporter; +use prometheus::{Encoder, TextEncoder}; +use reqwest::{header::AUTHORIZATION, Client, Method, StatusCode as HttpStatus}; +use sandbox::micro::{ + MicroConfig, MicroExecuteRequest, MicroImage, MicroStartRequest, SandboxMicro, +}; +use sandbox::run::{RunConfig, RunRequest, SandboxRun}; +use sandbox::{ + AgentContext, AgentContextFile, AgentDispatchRequest, AgentDispatcher, AgentDispatcherConfig, + AgentFileContent, AgentKind, AgentParameters, SandboxConfig, SandboxError, SandboxFs, + SandboxWasm, WasmConfig, WasmInvocation, WasmModuleSource, WasmValue, +}; +use serde::{Deserialize, Serialize}; +use serde_json::{json, Value}; +use sha2::{Digest, Sha256}; +use sqlx::postgres::PgPoolOptions; +use sqlx::types::Json; +use sqlx::{Error as SqlxError, PgPool, Row}; +use tower::ServiceBuilder; +use tower_http::cors::CorsLayer; +use tower_http::trace::TraceLayer; +use tracing::{dispatcher, error, info, warn}; +use tracing_opentelemetry::OpenTelemetryLayer; +use tracing_subscriber::prelude::*; +use uuid::Uuid; + +struct AppMetrics { + exporter: PrometheusExporter, + request_counter: Counter, + request_duration: Histogram, + sandbox_counter: Counter, + active_sessions: UpDownCounter, +} + +impl AppMetrics { + fn new() -> anyhow::Result { + let exporter = opentelemetry_prometheus::exporter().init()?; + let meter = global::meter("cyberdevstudio.api"); + let request_counter = meter + .u64_counter("api_requests_total") + .with_description("Total JSON-RPC requests processed by the API gateway") + .init(); + let request_duration = meter + .f64_histogram("api_request_duration_seconds") + .with_description("Latency of JSON-RPC method execution in seconds") + .init(); + let sandbox_counter = meter + .u64_counter("sandbox_operations_total") + .with_description("Sandbox operations executed by engine and operation") + .init(); + let active_sessions = meter + .i64_up_down_counter("active_sessions") + .with_description("Concurrent JSON-RPC requests being processed") + .init(); + Ok(Self { + exporter, + request_counter, + request_duration, + sandbox_counter, + active_sessions, + }) + } + + fn session_started(&self) { + self.active_sessions.add(1, &[]); + } + + fn session_finished(&self) { + self.active_sessions.add(-1, &[]); + } + + fn record_request( + &self, + method: &str, + status: &str, + duration: Duration, + role: &str, + auth_source: &str, + error_code: Option, + ) { + let mut attributes = vec![ + KeyValue::new("method", method.to_string()), + KeyValue::new("status", status.to_string()), + KeyValue::new("role", role.to_string()), + KeyValue::new("auth_source", auth_source.to_string()), + ]; + if let Some(code) = error_code { + attributes.push(KeyValue::new("error_code", code.to_string())); + } + self.request_counter.add(1, &attributes); + self.request_duration + .record(duration.as_secs_f64(), &attributes); + } + + fn record_sandbox_op(&self, engine: &'static str, operation: &'static str, success: bool) { + let status = if success { "success" } else { "error" }; + self.sandbox_counter.add( + 1, + &[ + KeyValue::new("engine", engine), + KeyValue::new("operation", operation), + KeyValue::new("status", status), + ], + ); + } + + fn render(&self) -> anyhow::Result { + let metric_families = self.exporter.registry().gather(); + let mut buffer = Vec::new(); + TextEncoder::new().encode(&metric_families, &mut buffer)?; + Ok(String::from_utf8(buffer)?) + } +} + +#[derive(Clone)] +struct AppState { + sandbox: Arc, + run: Arc, + wasm: Arc, + micro: Arc, + agents: Arc, + pool: PgPool, + auth: JwtVerifier, + llm: LlmClient, + metrics: Arc, +} + +#[derive(Clone)] +struct JwtVerifier { + decoding: DecodingKey, + validation: Validation, +} + +impl JwtVerifier { + fn from_env() -> anyhow::Result { + let secret = std::env::var("API_JWT_SECRET") + .or_else(|_| std::env::var("AUTH_JWT_SECRET")) + .map_err(|_| anyhow::anyhow!("API_JWT_SECRET environment variable is required"))?; + let issuer = + std::env::var("API_JWT_ISSUER").unwrap_or_else(|_| "cyber-dev-studio".to_string()); + let mut validation = Validation::new(Algorithm::HS256); + validation + .set_required_spec_claims(&["exp", "iat", "sub", "iss"]) + .expect("required claim configuration"); + validation.iss = Some(issuer); + Ok(Self { + decoding: DecodingKey::from_secret(secret.as_bytes()), + validation, + }) + } + + fn verify(&self, token: &str) -> std::result::Result { + decode::(token, &self.decoding, &self.validation) + .map(|data| data.claims) + .map_err(|_| RpcMethodError::unauthorized("invalid token")) + } +} + +#[derive(Debug, Deserialize)] +struct Claims { + sub: i32, + username: String, + role: String, + exp: usize, + iat: usize, + iss: String, + jti: String, +} + +#[derive(Debug, Clone)] +struct RequestContext { + user_id: i32, + username: String, + role: Role, + token_balance: i64, + api_key_id: Option, +} + +impl RequestContext { + fn require(&self, permission: Permission) -> std::result::Result<(), RpcMethodError> { + if self.role.allows(permission) { + Ok(()) + } else { + Err(RpcMethodError::forbidden("insufficient permissions")) + } + } + + fn auth_source(&self) -> &'static str { + if self.api_key_id.is_some() { + "api_key" + } else { + "jwt" + } + } + + fn ensure_tokens(&self) -> std::result::Result<(), RpcMethodError> { + if self.token_balance > 0 || self.is_admin() { + Ok(()) + } else { + Err(RpcMethodError::new( + -32092, + "insufficient token balance", + Some(json!({ "detail": "recharge required" })), + )) + } + } + + fn is_admin(&self) -> bool { + matches!(self.role, Role::Admin) + } +} + +#[derive(Debug, Clone, Copy, PartialEq, Eq)] +enum Role { + Admin, + Developer, + Viewer, +} + +impl Role { + fn parse(value: &str) -> Option { + match value { + "admin" => Some(Role::Admin), + "developer" => Some(Role::Developer), + "viewer" => Some(Role::Viewer), + _ => None, + } + } + + fn allows(self, permission: Permission) -> bool { + match permission { + Permission::FsRead | Permission::AgentView => true, + Permission::FsWrite + | Permission::Execute + | Permission::AgentControl + | Permission::LlmUse => matches!(self, Role::Admin | Role::Developer), + Permission::LlmAdmin => matches!(self, Role::Admin), + } + } + + fn as_str(self) -> &'static str { + match self { + Role::Admin => "admin", + Role::Developer => "developer", + Role::Viewer => "viewer", + } + } +} + +#[derive(Debug, Clone, Copy)] +enum Permission { + FsRead, + FsWrite, + Execute, + AgentView, + AgentControl, + LlmUse, + LlmAdmin, +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + init_tracing()?; + let metrics = Arc::new(AppMetrics::new()?); + let bind_addr = resolve_bind_address()?; + let pool = build_pool().await?; + let auth = JwtVerifier::from_env()?; + let (fs_sandbox, run_sandbox, wasm_sandbox, micro_sandbox) = initialize_sandboxes()?; + let agent_dispatcher = initialize_agent_dispatcher()?; + let llm = LlmClient::from_env()?; + + let sandbox = Arc::new(fs_sandbox); + let run = Arc::new(run_sandbox); + let wasm = Arc::new(wasm_sandbox); + let micro = Arc::new(micro_sandbox); + let agents = Arc::new(agent_dispatcher); + + let state = AppState { + sandbox, + run, + wasm, + micro, + agents, + pool, + auth, + llm, + metrics, + }; + + let app = Router::new() + .route("/health", get(health)) + .route("/metrics", get(metrics_endpoint)) + .route("/rpc", post(handle_rpc)) + .with_state(state) + .layer( + ServiceBuilder::new() + .layer(TraceLayer::new_for_http()) + .layer(CorsLayer::permissive()), + ); + + info!("binding", %bind_addr, "server starting"); + axum::Server::bind(&bind_addr) + .serve(app.into_make_service()) + .await?; + opentelemetry::global::shutdown_tracer_provider(); + Ok(()) +} + +fn init_tracing() -> anyhow::Result<()> { + if dispatcher::has_been_set() { + return Ok(()); + } + + let env_filter = tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| "info,tower_http=info".into()); + let fmt_layer = tracing_subscriber::fmt::layer() + .json() + .with_timer(tracing_subscriber::fmt::time::UtcTime::rfc3339()); + let registry = tracing_subscriber::registry() + .with(env_filter) + .with(fmt_layer); + + let otlp_endpoint = std::env::var("OTEL_EXPORTER_OTLP_TRACES_ENDPOINT") + .or_else(|_| std::env::var("OTEL_EXPORTER_OTLP_ENDPOINT")) + .ok(); + + if let Some(endpoint) = otlp_endpoint { + let resource = Resource::new(vec![ + KeyValue::new("service.name", "cyberdevstudio-api"), + KeyValue::new("service.namespace", "cyberdevstudio"), + KeyValue::new("service.version", env!("CARGO_PKG_VERSION")), + ]); + let exporter = opentelemetry_otlp::new_exporter() + .tonic() + .with_endpoint(endpoint); + let tracer = opentelemetry_otlp::new_pipeline() + .tracing() + .with_trace_config(Config::default().with_resource(resource)) + .with_exporter(exporter) + .install_batch(opentelemetry::runtime::Tokio)?; + registry.with(OpenTelemetryLayer::new(tracer)).try_init()?; + } else { + registry.try_init()?; + } + + Ok(()) +} + +fn resolve_bind_address() -> anyhow::Result { + let raw = std::env::var("API_BIND_ADDR").unwrap_or_else(|_| "0.0.0.0:6813".to_string()); + Ok(raw.parse()?) +} + +async fn build_pool() -> anyhow::Result { + let database_url = std::env::var("DATABASE_URL") + .map_err(|_| anyhow::anyhow!("DATABASE_URL environment variable is required"))?; + let max_connections = std::env::var("API_DATABASE_MAX_CONNECTIONS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(10); + let pool = PgPoolOptions::new() + .max_connections(max_connections) + .acquire_timeout(Duration::from_secs(10)) + .connect(&database_url) + .await?; + Ok(pool) +} + +fn initialize_sandboxes() -> anyhow::Result<(SandboxFs, SandboxRun, SandboxWasm, SandboxMicro)> { + let max_size = std::env::var("SANDBOX_MAX_FILE_SIZE") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(512 * 1024); + let root = sandbox_root()?; + + let fs = SandboxFs::new(SandboxConfig::new(root.clone(), max_size)?); + + let allowed_programs = std::env::var("SANDBOX_RUN_ALLOWED") + .ok() + .map(|value| { + value + .split(',') + .map(|item| item.trim().to_string()) + .filter(|item| !item.is_empty()) + .collect::>() + }) + .filter(|items| !items.is_empty()) + .unwrap_or_else(|| vec!["/bin/sh".to_string(), "/usr/bin/env".to_string()]); + + let env_allowlist = std::env::var("SANDBOX_RUN_ENV_ALLOW") + .ok() + .map(|value| { + value + .split(',') + .map(|item| item.trim().to_string()) + .filter(|item| !item.is_empty()) + .collect::>() + }) + .unwrap_or_else(|| vec!["PATH".to_string()]); + + let path_env = + std::env::var("SANDBOX_RUN_PATH").unwrap_or_else(|_| "/usr/bin:/bin".to_string()); + let mut fixed_env = vec![ + ("PATH".to_string(), path_env), + ("HOME".to_string(), root.to_string_lossy().to_string()), + ]; + + if let Ok(extra_fixed) = std::env::var("SANDBOX_RUN_FIXED_ENV") { + for pair in extra_fixed + .split(',') + .map(|p| p.trim()) + .filter(|p| !p.is_empty()) + { + if let Some((key, value)) = pair.split_once('=') { + fixed_env.push((key.trim().to_string(), value.trim().to_string())); + } + } + } + + let default_timeout_ms = std::env::var("SANDBOX_RUN_DEFAULT_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(10_000); + let max_timeout_ms = std::env::var("SANDBOX_RUN_MAX_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(30_000); + let max_output_bytes_raw = std::env::var("SANDBOX_RUN_MAX_OUTPUT_BYTES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(512 * 1024); + let max_output_bytes = usize::try_from(max_output_bytes_raw) + .map_err(|_| anyhow::anyhow!("SANDBOX_RUN_MAX_OUTPUT_BYTES exceeds platform limits"))?; + + let run_config = RunConfig::new( + &root, + allowed_programs, + env_allowlist, + fixed_env, + Duration::from_millis(default_timeout_ms), + Duration::from_millis(max_timeout_ms), + max_output_bytes, + )?; + + let wasm_memory_limit = std::env::var("SANDBOX_WASM_MAX_MEMORY_BYTES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(64 * 1024 * 1024); + let wasm_table_limit = std::env::var("SANDBOX_WASM_MAX_TABLE_ELEMENTS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(2_048); + let wasm_default_fuel = std::env::var("SANDBOX_WASM_DEFAULT_FUEL") + .ok() + .and_then(|v| v.parse::().ok()); + + let wasm_config = WasmConfig::new( + root.clone(), + wasm_memory_limit, + wasm_table_limit, + wasm_default_fuel, + )?; + + let micro_default_timeout_ms = std::env::var("SANDBOX_MICRO_DEFAULT_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(5_000); + let micro_max_timeout_ms = std::env::var("SANDBOX_MICRO_MAX_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(30_000); + let micro_max_output_bytes_raw = std::env::var("SANDBOX_MICRO_MAX_OUTPUT_BYTES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(256 * 1024); + let micro_max_output_bytes = usize::try_from(micro_max_output_bytes_raw) + .map_err(|_| anyhow::anyhow!("SANDBOX_MICRO_MAX_OUTPUT_BYTES exceeds platform limits"))?; + + let micro_images = resolve_micro_images()?; + let micro_base_env = resolve_micro_base_env(); + let micro_config = MicroConfig::new( + &root, + micro_images, + Duration::from_millis(micro_default_timeout_ms), + Duration::from_millis(micro_max_timeout_ms), + micro_max_output_bytes, + micro_base_env, + )?; + + Ok(( + fs, + SandboxRun::new(run_config), + SandboxWasm::new(wasm_config), + SandboxMicro::new(micro_config), + )) +} + +fn initialize_agent_dispatcher() -> anyhow::Result { + let endpoint = + std::env::var("AGENT_LLM_ENDPOINT").unwrap_or_else(|_| "http://localhost:6988".to_string()); + let default_model = + std::env::var("AGENT_DEFAULT_MODEL").unwrap_or_else(|_| "nous-hermes-2-3b.Q4".to_string()); + let timeout_ms = std::env::var("AGENT_LLM_TIMEOUT_MS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(30_000); + let history_capacity = std::env::var("AGENT_HISTORY_CAPACITY") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(128); + let context_limit = std::env::var("AGENT_CONTEXT_LIMIT_BYTES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(512 * 1024); + let api_key = std::env::var("AGENT_LLM_API_KEY").ok(); + + let config = AgentDispatcherConfig::new(endpoint, default_model) + .with_timeout(Duration::from_millis(timeout_ms)) + .with_history_capacity(history_capacity) + .with_context_limit(context_limit) + .with_api_key(api_key); + + AgentDispatcher::new(config).map_err(|err| anyhow::anyhow!(err.to_string())) +} + +fn sandbox_root() -> anyhow::Result { + let raw = std::env::var("SANDBOX_ROOT").unwrap_or_else(|_| "./data/sandbox".to_string()); + let path = PathBuf::from(&raw); + if path.is_absolute() { + Ok(path) + } else { + let cwd = std::env::current_dir()?; + Ok(cwd.join(path)) + } +} + +fn resolve_micro_images() -> anyhow::Result> { + if let Ok(raw) = std::env::var("SANDBOX_MICRO_IMAGES") { + let definitions: Vec = serde_json::from_str(&raw) + .map_err(|err| anyhow::anyhow!("failed to parse SANDBOX_MICRO_IMAGES: {err}"))?; + if definitions.is_empty() { + anyhow::bail!("SANDBOX_MICRO_IMAGES must define at least one image"); + } + let mut images = Vec::with_capacity(definitions.len()); + for definition in definitions { + let extension = definition + .extension + .unwrap_or_else(|| guess_extension(&definition.name).to_string()); + let env_pairs = definition + .env + .into_iter() + .map(|pair| (pair.key, pair.value)) + .collect::>(); + images.push(MicroImage::new( + definition.name, + definition.command, + definition.args, + extension, + env_pairs, + )?); + } + Ok(images) + } else { + default_micro_images() + } +} + +fn default_micro_images() -> anyhow::Result> { + let python_command = std::env::var("SANDBOX_MICRO_PYTHON") + .ok() + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| detect_binary("python3").unwrap_or_else(|| "python3".to_string())); + let node_command = std::env::var("SANDBOX_MICRO_NODE") + .ok() + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| detect_binary("node").unwrap_or_else(|| "node".to_string())); + + let mut images = Vec::new(); + images.push(MicroImage::new( + "python", + python_command, + vec!["-u".to_string()], + "py", + vec![("PYTHONUNBUFFERED".to_string(), "1".to_string())], + )?); + images.push(MicroImage::new( + "node", + node_command, + Vec::new(), + "js", + Vec::new(), + )?); + Ok(images) +} + +fn detect_binary(name: &str) -> Option { + let path = std::env::var("PATH").ok()?; + for entry in path + .split(':') + .map(|segment| segment.trim()) + .filter(|s| !s.is_empty()) + { + let candidate = Path::new(entry).join(name); + if let Ok(metadata) = std::fs::metadata(&candidate) { + if metadata.is_file() { + return Some(candidate.to_string_lossy().to_string()); + } + } + } + None +} + +fn guess_extension(name: &str) -> &'static str { + let lower = name.to_ascii_lowercase(); + if lower.contains("python") { + "py" + } else if lower.contains("node") || lower.contains("js") { + "js" + } else if lower.contains("ruby") { + "rb" + } else if lower.contains("go") { + "go" + } else { + "txt" + } +} + +fn resolve_micro_base_env() -> Vec<(String, String)> { + let path_env = std::env::var("SANDBOX_MICRO_PATH") + .ok() + .filter(|value| !value.trim().is_empty()) + .unwrap_or_else(|| std::env::var("PATH").unwrap_or_else(|_| "/usr/bin:/bin".to_string())); + let mut base = vec![ + ("PATH".to_string(), path_env), + ("LANG".to_string(), "C".to_string()), + ("LC_ALL".to_string(), "C".to_string()), + ("TERM".to_string(), "dumb".to_string()), + ]; + if let Ok(extra) = std::env::var("SANDBOX_MICRO_BASE_ENV") { + for pair in extra + .split(',') + .map(|segment| segment.trim()) + .filter(|segment| !segment.is_empty()) + { + if let Some((key, value)) = pair.split_once('=') { + base.push((key.trim().to_string(), value.trim().to_string())); + } + } + } + base +} + +async fn health() -> impl IntoResponse { + (StatusCode::OK, Json(json!({ "status": "ok" }))) +} + +async fn metrics_endpoint(State(state): State) -> Response { + match state.metrics.render() { + Ok(body) => (StatusCode::OK, body).into_response(), + Err(err) => { + error!(?err, "failed to encode metrics"); + ( + StatusCode::INTERNAL_SERVER_ERROR, + "failed to encode metrics", + ) + .into_response() + } + } +} + +async fn authenticate_request( + state: &AppState, + headers: &HeaderMap, +) -> std::result::Result { + if let Some(value) = headers.get("x-api-key") { + if !value.as_bytes().is_empty() { + return authenticate_with_api_key(state, value).await; + } + } + + let authorization = headers + .get(axum::http::header::AUTHORIZATION) + .ok_or_else(|| RpcMethodError::unauthorized("missing authorization header"))?; + let authorization = authorization + .to_str() + .map_err(|_| RpcMethodError::unauthorized("invalid authorization header"))?; + let token = authorization + .strip_prefix("Bearer ") + .ok_or_else(|| RpcMethodError::unauthorized("unsupported authorization scheme"))?; + authenticate_with_jwt(state, token).await +} + +async fn authenticate_with_api_key( + state: &AppState, + value: &axum::http::HeaderValue, +) -> std::result::Result { + let api_key = value + .to_str() + .map_err(|_| RpcMethodError::unauthorized("invalid api key header"))?; + if api_key.is_empty() { + return Err(RpcMethodError::unauthorized("invalid api key")); + } + let hash = hash_api_key(api_key); + let row = sqlx::query( + "SELECT api_keys.id AS api_key_id, users.id AS user_id, users.username, users.role, users.token_balance \ + FROM api_keys JOIN users ON users.id = api_keys.user_id WHERE api_keys.api_key_hash = $1", + ) + .bind(&hash) + .fetch_optional(&state.pool) + .await + .map_err(|err| RpcMethodError::internal(&err.to_string()))?; + + let row = row.ok_or_else(|| RpcMethodError::unauthorized("invalid api key"))?; + let role_str: String = row.get("role"); + let role = Role::parse(&role_str) + .ok_or_else(|| RpcMethodError::internal("user has unsupported role"))?; + + let api_key_id: Uuid = row.get("api_key_id"); + let context = RequestContext { + user_id: row.get("user_id"), + username: row.get("username"), + role, + token_balance: row.get("token_balance"), + api_key_id: Some(api_key_id), + }; + + if let Err(err) = sqlx::query("UPDATE api_keys SET last_used_at = NOW() WHERE id = $1") + .bind(api_key_id) + .execute(&state.pool) + .await + { + warn!("failed to update api key usage", error = %err); + } + + Ok(context) +} + +async fn authenticate_with_jwt( + state: &AppState, + token: &str, +) -> std::result::Result { + let claims = state.auth.verify(token)?; + let row = sqlx::query("SELECT username, role, token_balance FROM users WHERE id = $1") + .bind(claims.sub) + .fetch_one(&state.pool) + .await + .map_err(|err| match err { + sqlx::Error::RowNotFound => RpcMethodError::unauthorized("user not found"), + other => RpcMethodError::internal(&other.to_string()), + })?; + + let role_str: String = row.get("role"); + let role = Role::parse(&role_str) + .ok_or_else(|| RpcMethodError::internal("user has unsupported role"))?; + + Ok(RequestContext { + user_id: claims.sub, + username: row.get("username"), + role, + token_balance: row.get("token_balance"), + api_key_id: None, + }) +} + +fn hash_api_key(key: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(key.as_bytes()); + hex_encode(hasher.finalize()) +} + +async fn handle_rpc( + State(state): State, + headers: HeaderMap, + Json(req): Json, +) -> impl IntoResponse { + let method_label = req.method.clone(); + state.metrics.session_started(); + let start = Instant::now(); + + if req.jsonrpc != "2.0" { + let response = Json(RpcResponse::error( + req.id, + -32600, + "invalid jsonrpc version", + None, + )); + state.metrics.record_request( + &method_label, + "invalid", + start.elapsed(), + "unknown", + "unauthenticated", + Some(-32600), + ); + state.metrics.session_finished(); + return response; + } + + let mut role_label = "unknown"; + let mut auth_source = "unauthenticated"; + let ctx = match authenticate_request(&state, &headers).await { + Ok(ctx) => { + role_label = ctx.role.as_str(); + auth_source = ctx.auth_source(); + ctx + } + Err(err) => { + error!("authentication failed", message = %err.message); + let response = Json(RpcResponse::error(req.id, err.code, &err.message, err.data)); + state.metrics.record_request( + &method_label, + "unauthorized", + start.elapsed(), + role_label, + auth_source, + Some(err.code), + ); + state.metrics.session_finished(); + return response; + } + }; + + let (response, status, error_code) = + match process_request(&state, &ctx, req.method, req.params).await { + Ok(result) => (Json(RpcResponse::success(req.id, result)), "ok", None), + Err(err) => { + error!("rpc error", message = %err.message); + ( + Json(RpcResponse::error(req.id, err.code, &err.message, err.data)), + "error", + Some(err.code), + ) + } + }; + state.metrics.record_request( + &method_label, + status, + start.elapsed(), + role_label, + auth_source, + error_code, + ); + state.metrics.session_finished(); + response +} + +async fn process_request( + state: &AppState, + ctx: &RequestContext, + method: String, + params: Option, +) -> std::result::Result { + match method.as_str() { + "fs.read" => { + ctx.require(Permission::FsRead)?; + let params: FsPathParams = parse_params(params)?; + let bytes = state.sandbox.read(Path::new(¶ms.path)).map_err(|err| { + state.metrics.record_sandbox_op("fs", "read", false); + RpcMethodError::from_sandbox(-32001, "failed to read file", err) + })?; + state.metrics.record_sandbox_op("fs", "read", true); + Ok(json!({ "data": BASE64.encode(bytes) })) + } + "fs.write" => { + ctx.require(Permission::FsWrite)?; + let params: FsWriteParams = parse_params(params)?; + let data = BASE64.decode(params.data.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + state + .sandbox + .write(Path::new(¶ms.path), data) + .map_err(|err| { + state.metrics.record_sandbox_op("fs", "write", false); + RpcMethodError::from_sandbox(-32002, "failed to write file", err) + })?; + state.metrics.record_sandbox_op("fs", "write", true); + Ok(json!({ "status": "ok" })) + } + "fs.list" => { + ctx.require(Permission::FsRead)?; + let params: FsPathParams = parse_params(params)?; + let entries = state.sandbox.list(Path::new(¶ms.path)).map_err(|err| { + state.metrics.record_sandbox_op("fs", "list", false); + RpcMethodError::from_sandbox(-32003, "failed to list directory", err) + })?; + state.metrics.record_sandbox_op("fs", "list", true); + Ok(serde_json::to_value(entries).expect("serialize entries")) + } + "fs.delete" => { + ctx.require(Permission::FsWrite)?; + let params: FsPathParams = parse_params(params)?; + state + .sandbox + .delete(Path::new(¶ms.path)) + .map_err(|err| { + state.metrics.record_sandbox_op("fs", "delete", false); + RpcMethodError::from_sandbox(-32004, "failed to delete path", err) + })?; + state.metrics.record_sandbox_op("fs", "delete", true); + Ok(json!({ "status": "ok" })) + } + "fs.mkdir" => { + ctx.require(Permission::FsWrite)?; + let params: FsPathParams = parse_params(params)?; + state + .sandbox + .mkdir(Path::new(¶ms.path)) + .map_err(|err| { + state.metrics.record_sandbox_op("fs", "mkdir", false); + RpcMethodError::from_sandbox(-32005, "failed to create directory", err) + })?; + state.metrics.record_sandbox_op("fs", "mkdir", true); + Ok(json!({ "status": "ok" })) + } + "project.create" => { + ctx.require(Permission::FsWrite)?; + let params: ProjectCreateParams = parse_params(params)?; + let name = normalize_project_name(¶ms.name)?; + let description = params.description.as_ref().map(|d| truncate_description(d)); + let record = create_project(&state.pool, ctx, &name, description.as_deref()).await?; + let project_root = project_directory_relative(&record.id); + state.sandbox.mkdir(&project_root).map_err(|err| { + state.metrics.record_sandbox_op("fs", "mkdir", false); + RpcMethodError::from_sandbox(-32050, "failed to prepare project", err) + })?; + state.metrics.record_sandbox_op("fs", "mkdir", true); + let activity_name = record.name.clone(); + record_project_activity( + &state.pool, + record.id, + ctx.user_id, + "project.created", + Some(json!({ "name": activity_name })), + ) + .await + .map_err(|err| map_db_activity_error(err, "failed to record project activity"))?; + Ok(record.to_value()) + } + "project.list" => { + ctx.require(Permission::FsRead)?; + let projects = list_projects(&state.pool, ctx).await?; + Ok(Value::Array(projects)) + } + "project.open" => { + ctx.require(Permission::FsRead)?; + let params: ProjectOpenParams = parse_params(params)?; + let project_id = parse_project_id(¶ms.project_id)?; + let record = load_project(&state.pool, ctx, &project_id).await?; + let include_content = params.include_content.unwrap_or(false); + let files = project_files(&state.pool, &project_id, include_content).await?; + Ok(json!({ + "project": record.to_value(), + "files": files, + })) + } + "project.delete" => { + ctx.require(Permission::FsWrite)?; + let params: ProjectIdParams = parse_params(params)?; + let project_id = parse_project_id(¶ms.project_id)?; + let record = load_project(&state.pool, ctx, &project_id).await?; + delete_project(&state.pool, &project_id).await?; + let project_root = project_directory_relative(&project_id); + state.sandbox.delete(&project_root).map_err(|err| { + state.metrics.record_sandbox_op("fs", "delete", false); + RpcMethodError::from_sandbox(-32054, "failed to remove project files", err) + })?; + state.metrics.record_sandbox_op("fs", "delete", true); + let name = record.name.clone(); + record_project_activity( + &state.pool, + project_id, + ctx.user_id, + "project.deleted", + Some(json!({ "name": name })), + ) + .await + .map_err(|err| map_db_activity_error(err, "failed to record project activity"))?; + Ok(json!({ "status": "ok" })) + } + "project.file.save" => { + ctx.require(Permission::FsWrite)?; + let params: ProjectFileSaveParams = parse_params(params)?; + let project_id = parse_project_id(¶ms.project_id)?; + let _ = load_project(&state.pool, ctx, &project_id).await?; + let encoding = params.encoding.unwrap_or_else(|| "base64".to_string()); + if encoding.to_lowercase() != "base64" { + return Err(RpcMethodError::new( + -32602, + "unsupported file encoding", + Some(json!({ "detail": encoding })), + )); + } + let data = BASE64.decode(params.data.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let relative_path = normalize_project_path(¶ms.path)?; + let sha256 = Sha256::digest(&data); + let saved = + save_project_file(&state.pool, &project_id, &relative_path, &data, &sha256).await?; + let project_root = project_directory_relative(&project_id).join(&relative_path); + state.sandbox.write(project_root, &data).map_err(|err| { + state.metrics.record_sandbox_op("fs", "write", false); + RpcMethodError::from_sandbox(-32051, "failed to persist project file", err) + })?; + state.metrics.record_sandbox_op("fs", "write", true); + if let Some(message) = params.message { + if !message.trim().is_empty() { + record_project_activity( + &state.pool, + project_id, + ctx.user_id, + "project.file.save", + Some(json!({ + "path": relative_path.to_string_lossy(), + "message": message.trim(), + })), + ) + .await + .map_err(|err| { + map_db_activity_error(err, "failed to record project activity") + })?; + } + } else { + record_project_activity( + &state.pool, + project_id, + ctx.user_id, + "project.file.save", + Some(json!({ + "path": relative_path.to_string_lossy(), + })), + ) + .await + .map_err(|err| map_db_activity_error(err, "failed to record project activity"))?; + } + Ok(saved) + } + "project.file.read" => { + ctx.require(Permission::FsRead)?; + let params: ProjectFilePathParams = parse_params(params)?; + let project_id = parse_project_id(¶ms.project_id)?; + let _ = load_project(&state.pool, ctx, &project_id).await?; + let relative_path = normalize_project_path(¶ms.path)?; + let file = read_project_file(&state.pool, &project_id, &relative_path).await?; + Ok(file) + } + "project.file.delete" => { + ctx.require(Permission::FsWrite)?; + let params: ProjectFilePathParams = parse_params(params)?; + let project_id = parse_project_id(¶ms.project_id)?; + let _ = load_project(&state.pool, ctx, &project_id).await?; + let relative_path = normalize_project_path(¶ms.path)?; + delete_project_file(&state.pool, &project_id, &relative_path).await?; + let project_root = project_directory_relative(&project_id).join(&relative_path); + state.sandbox.delete(project_root).map_err(|err| { + state.metrics.record_sandbox_op("fs", "delete", false); + RpcMethodError::from_sandbox(-32053, "failed to delete project file", err) + })?; + state.metrics.record_sandbox_op("fs", "delete", true); + record_project_activity( + &state.pool, + project_id, + ctx.user_id, + "project.file.delete", + Some(json!({ "path": relative_path.to_string_lossy() })), + ) + .await + .map_err(|err| map_db_activity_error(err, "failed to record project activity"))?; + Ok(json!({ "status": "ok" })) + } + "run.exec" => { + ctx.require(Permission::Execute)?; + let params: RunExecParams = parse_params(params)?; + let request = params.into_request()?; + let result = state.run.execute(request).await.map_err(|err| { + state.metrics.record_sandbox_op("run", "exec", false); + RpcMethodError::from_sandbox(-32010, "failed to execute process", err) + })?; + state.metrics.record_sandbox_op("run", "exec", true); + Ok(json!({ + "exit_code": result.exit_code, + "stdout": BASE64.encode(result.stdout), + "stderr": BASE64.encode(result.stderr), + "duration_ms": result.duration.as_millis() + })) + } + "run.describe" => { + ctx.require(Permission::FsRead)?; + let config = state.run.config(); + let allowed: Vec = config.allowed_programs().cloned().collect(); + Ok(json!({ + "root": config.root().display().to_string(), + "allowed_programs": allowed, + "default_timeout_ms": config.default_timeout().as_millis(), + "max_timeout_ms": config.max_timeout().as_millis(), + "max_output_bytes": config.max_output_bytes() + })) + } + "wasm.invoke" => { + ctx.require(Permission::Execute)?; + let params: WasmInvokeParams = parse_params(params)?; + let module_source = resolve_wasm_module(¶ms)?; + let wasm_params = params + .params + .into_iter() + .map(WasmParam::into_value) + .collect::, _>>() + .map_err(|err| RpcMethodError::new(-32602, err.as_str(), None))?; + + let mut invocation = + WasmInvocation::new(module_source, params.function).with_params(wasm_params); + if let Some(fuel) = params.fuel { + invocation = invocation.with_fuel(fuel); + } + if let Some(memory) = params.memory_limit { + invocation = invocation.with_memory_limit(memory); + } + if let Some(table) = params.table_elements_limit { + invocation = invocation.with_table_elements_limit(table); + } + + let values = state.wasm.invoke(invocation).map_err(|err| { + state.metrics.record_sandbox_op("wasm", "invoke", false); + RpcMethodError::from_sandbox(-32020, "failed to execute wasm", err) + })?; + state.metrics.record_sandbox_op("wasm", "invoke", true); + let serialized: Vec = values.into_iter().map(wasm_value_to_json).collect(); + Ok(json!({ "values": serialized })) + } + "wasm.describe" => { + ctx.require(Permission::FsRead)?; + let config = state.wasm.config(); + Ok(json!({ + "root": config.root().display().to_string(), + "max_memory_bytes": config.max_memory_bytes(), + "max_table_elements": config.max_table_elements(), + "default_fuel": config.default_fuel(), + })) + } + "micro.start" => { + ctx.require(Permission::Execute)?; + let params: MicroStartParams = parse_params(params)?; + let init_script = match params.init_script { + Some(ref value) if !value.is_empty() => { + let bytes = BASE64.decode(value.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + Some(String::from_utf8(bytes).map_err(|err| { + RpcMethodError::new( + -32602, + "init script must be valid utf-8", + Some(json!({ "detail": err.to_string() })), + ) + })?) + } + _ => None, + }; + let request = MicroStartRequest { + image: params.image, + init_script, + }; + let instance = state.micro.start(request).await.map_err(|err| { + state.metrics.record_sandbox_op("micro", "start", false); + RpcMethodError::from_sandbox(-32030, "failed to start micro vm", err) + })?; + state.metrics.record_sandbox_op("micro", "start", true); + Ok(json!({ + "vm_id": instance.id().to_string(), + "image": instance.image().to_string(), + "working_dir": instance.workdir().display().to_string(), + })) + } + "micro.execute" => { + ctx.require(Permission::Execute)?; + let params: MicroExecuteParams = parse_params(params)?; + let vm_id = Uuid::parse_str(¶ms.vm_id).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid vm identifier", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let code_bytes = BASE64.decode(params.code.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let code = String::from_utf8(code_bytes).map_err(|err| { + RpcMethodError::new( + -32602, + "code must be valid utf-8", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let request = MicroExecuteRequest { + vm_id, + code, + timeout: params.timeout_ms.map(Duration::from_millis), + }; + let result = state.micro.execute(request).await.map_err(|err| { + state.metrics.record_sandbox_op("micro", "execute", false); + RpcMethodError::from_sandbox(-32031, "failed to execute micro vm code", err) + })?; + state.metrics.record_sandbox_op("micro", "execute", true); + Ok(json!({ + "exit_code": result.exit_code, + "stdout": BASE64.encode(result.stdout), + "stderr": BASE64.encode(result.stderr), + "duration_ms": result.duration.as_millis(), + })) + } + "micro.stop" => { + ctx.require(Permission::Execute)?; + let params: MicroStopParams = parse_params(params)?; + let vm_id = Uuid::parse_str(¶ms.vm_id).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid vm identifier", + Some(json!({ "detail": err.to_string() })), + ) + })?; + state.micro.stop(vm_id).await.map_err(|err| { + state.metrics.record_sandbox_op("micro", "stop", false); + RpcMethodError::from_sandbox(-32032, "failed to stop micro vm", err) + })?; + state.metrics.record_sandbox_op("micro", "stop", true); + Ok(json!({ "status": "ok" })) + } + "micro.describe" => { + ctx.require(Permission::FsRead)?; + let config = state.micro.config(); + let images: Vec = config + .images() + .map(|image| { + json!({ + "name": image.name(), + "command": image.command(), + "args": image.args().cloned().collect::>(), + "extension": image.extension(), + "env": image + .env() + .map(|(key, value)| json!({ "key": key, "value": value })) + .collect::>(), + }) + }) + .collect(); + let base_env: Vec = config + .base_env() + .iter() + .map(|(key, value)| json!({ "key": key, "value": value })) + .collect(); + Ok(json!({ + "root": config.root().display().to_string(), + "default_timeout_ms": config.default_timeout().as_millis(), + "max_timeout_ms": config.max_timeout().as_millis(), + "max_output_bytes": config.max_output_bytes(), + "images": images, + "base_env": base_env, + })) + } + "llm.chat" => { + ctx.require(Permission::LlmUse)?; + ctx.ensure_tokens()?; + let params: LlmChatParams = parse_params(params)?; + state.llm.chat(ctx, params).await + } + "llm.completion" | "llm.completions" => { + ctx.require(Permission::LlmUse)?; + ctx.ensure_tokens()?; + let params: LlmCompletionParams = parse_params(params)?; + state.llm.completion(ctx, params).await + } + "llm.embed" => { + ctx.require(Permission::LlmUse)?; + ctx.ensure_tokens()?; + let params: LlmEmbedParams = parse_params(params)?; + state.llm.embed(ctx, params).await + } + "llm.list_models" => { + ctx.require(Permission::LlmAdmin)?; + state.llm.list_models().await + } + "llm.status" => { + ctx.require(Permission::LlmAdmin)?; + state.llm.status().await + } + "llm.download" => { + ctx.require(Permission::LlmAdmin)?; + let params: LlmModelParams = parse_params(params)?; + state.llm.download(ctx, ¶ms).await + } + "llm.start" => { + ctx.require(Permission::LlmAdmin)?; + let params: LlmAdminLoadParams = parse_params(params)?; + state.llm.load(ctx, params).await + } + "llm.stop" => { + ctx.require(Permission::LlmAdmin)?; + let params: LlmModelParams = parse_params(params)?; + state.llm.unload(ctx, ¶ms).await + } + "agent.list" => { + ctx.require(Permission::AgentView)?; + let agents = state.agents.list_agents(); + Ok(serde_json::to_value(agents).expect("serialize agents")) + } + "agent.history" => { + ctx.require(Permission::AgentView)?; + let params: AgentHistoryParams = parse_params(params)?; + let mut limit = params.limit.unwrap_or(20); + if limit == 0 { + limit = 1; + } + if limit > 256 { + limit = 256; + } + let history = state.agents.history(limit); + Ok(serde_json::to_value(history).expect("serialize history")) + } + "agent.status" => { + ctx.require(Permission::AgentView)?; + let params: AgentStatusParams = parse_params(params)?; + let task_id = Uuid::parse_str(¶ms.task_id).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid task identifier", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let snapshot = state + .agents + .status(&task_id) + .ok_or_else(|| RpcMethodError::new(-32041, "agent task not found", None))?; + Ok(serde_json::to_value(snapshot).expect("serialize status")) + } + "agent.cancel" => { + ctx.require(Permission::AgentControl)?; + let params: AgentStatusParams = parse_params(params)?; + let task_id = Uuid::parse_str(¶ms.task_id).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid task identifier", + Some(json!({ "detail": err.to_string() })), + ) + })?; + let snapshot = state.agents.cancel(&task_id).map_err(|err| { + RpcMethodError::from_sandbox(-32042, "failed to cancel agent", err) + })?; + Ok(serde_json::to_value(snapshot).expect("serialize status")) + } + "agent.dispatch" => { + ctx.require(Permission::AgentControl)?; + let params: AgentDispatchParams = parse_params(params)?; + let AgentDispatchParams { + agent, + objective, + context, + model, + metadata, + parameters, + } = params; + let context = build_agent_context(&state.sandbox, context).map_err(|err| { + RpcMethodError::from_sandbox(-32043, "failed to prepare agent context", err) + })?; + let parameters = parameters.map(AgentParameterOverrides::into_parameters); + let metadata = enrich_agent_metadata(metadata, ctx); + let request = AgentDispatchRequest { + agent, + objective, + context, + model, + metadata, + parameters, + }; + let submission = state.agents.dispatch(request).map_err(|err| { + RpcMethodError::from_sandbox(-32040, "failed to dispatch agent", err) + })?; + Ok(json!({ + "task_id": submission.id.to_string(), + "status": submission.status, + })) + } + _ => Err(RpcMethodError::new(-32601, "method not found", None)), + } +} + +#[derive(Clone)] +struct LlmClient { + http: Client, + base_url: String, + admin_token: Option, +} + +impl LlmClient { + fn from_env() -> anyhow::Result { + let base_url = + std::env::var("LLM_SERVER_URL").unwrap_or_else(|_| "http://127.0.0.1:6988".to_string()); + let admin_token = std::env::var("LLM_SERVER_ADMIN_TOKEN").ok(); + let timeout_secs = std::env::var("LLM_HTTP_TIMEOUT_SECS") + .ok() + .and_then(|value| value.parse::().ok()) + .unwrap_or(30); + let http = Client::builder() + .timeout(Duration::from_secs(timeout_secs)) + .build()?; + Ok(Self { + http, + base_url, + admin_token, + }) + } + + async fn chat( + &self, + ctx: &RequestContext, + params: LlmChatParams, + ) -> std::result::Result { + self.post_user("/v1/chat/completions", ¶ms, ctx).await + } + + async fn completion( + &self, + ctx: &RequestContext, + params: LlmCompletionParams, + ) -> std::result::Result { + self.post_user("/v1/completions", ¶ms, ctx).await + } + + async fn embed( + &self, + ctx: &RequestContext, + params: LlmEmbedParams, + ) -> std::result::Result { + self.post_user("/v1/embeddings", ¶ms, ctx).await + } + + async fn list_models(&self) -> std::result::Result { + self.get_admin("/admin/models").await + } + + async fn status(&self) -> std::result::Result { + self.get_admin("/admin/status").await + } + + async fn download( + &self, + ctx: &RequestContext, + params: &LlmModelParams, + ) -> std::result::Result { + self.post_admin("/admin/download", params, Some(ctx)).await + } + + async fn load( + &self, + ctx: &RequestContext, + params: LlmAdminLoadParams, + ) -> std::result::Result { + self.post_admin("/admin/load", ¶ms, Some(ctx)).await + } + + async fn unload( + &self, + ctx: &RequestContext, + params: &LlmModelParams, + ) -> std::result::Result { + self.post_admin("/admin/unload", params, Some(ctx)).await + } + + async fn post_user( + &self, + path: &str, + body: &T, + ctx: &RequestContext, + ) -> std::result::Result { + let request_id = Uuid::new_v4(); + self.send_request( + Method::POST, + path, + Some(body), + Some(ctx), + false, + Some(request_id), + ) + .await + } + + async fn post_admin( + &self, + path: &str, + body: &T, + ctx: Option<&RequestContext>, + ) -> std::result::Result { + self.send_request( + Method::POST, + path, + Some(body), + ctx, + true, + Some(Uuid::new_v4()), + ) + .await + } + + async fn get_admin(&self, path: &str) -> std::result::Result { + self.send_request::(Method::GET, path, None, None, true, Some(Uuid::new_v4())) + .await + } + + async fn send_request( + &self, + method: Method, + path: &str, + body: Option<&T>, + ctx: Option<&RequestContext>, + admin: bool, + request_id: Option, + ) -> std::result::Result { + let url = format!( + "{}/{}", + self.base_url.trim_end_matches('/'), + path.trim_start_matches('/') + ); + let mut builder = self.http.request(method, url); + if let Some(ctx) = ctx { + builder = builder.header("X-User-Id", ctx.user_id.to_string()).header( + "X-Request-Id", + request_id.unwrap_or_else(Uuid::new_v4).to_string(), + ); + } else if let Some(request_id) = request_id { + builder = builder.header("X-Request-Id", request_id.to_string()); + } + if admin { + let token = self + .admin_token + .as_ref() + .ok_or_else(|| RpcMethodError::internal("LLM_SERVER_ADMIN_TOKEN not configured"))?; + builder = builder.header(AUTHORIZATION, format!("Bearer {token}")); + } + if let Some(body) = body { + builder = builder.json(body); + } + let response = builder + .send() + .await + .map_err(|err| RpcMethodError::internal(&err.to_string()))?; + self.handle_response(response).await + } + + async fn handle_response( + &self, + response: reqwest::Response, + ) -> std::result::Result { + let status = response.status(); + let bytes = response + .bytes() + .await + .map_err(|err| RpcMethodError::internal(&err.to_string()))?; + let body: Value = serde_json::from_slice(&bytes).unwrap_or_else( + |_| json!({ "error": String::from_utf8_lossy(&bytes).trim().to_string() }), + ); + if status.is_success() { + return Ok(body); + } + let message = body + .get("error") + .and_then(|value| value.as_str()) + .filter(|value| !value.is_empty()) + .unwrap_or_else(|| status.canonical_reason().unwrap_or("request failed")); + let error = match status { + HttpStatus::UNAUTHORIZED => RpcMethodError::unauthorized(message), + HttpStatus::FORBIDDEN => RpcMethodError::forbidden(message), + HttpStatus::TOO_MANY_REQUESTS => RpcMethodError::new( + -32093, + "insufficient token balance", + Some(json!({ "detail": message })), + ), + HttpStatus::NOT_FOUND => RpcMethodError::new(-32044, message, Some(body.clone())), + _ => RpcMethodError::internal(message), + }; + Err(error) + } +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +struct LlmChatParams { + model: String, + messages: Vec, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_k: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + repeat_penalty: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + max_tokens: Option, +} + +#[derive(Debug, Deserialize, Serialize)] +struct LlmChatMessage { + role: String, + content: String, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +struct LlmCompletionParams { + model: String, + prompt: String, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_k: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + repeat_penalty: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + max_tokens: Option, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +struct LlmEmbedParams { + model: String, + input: LlmEmbedInput, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(untagged)] +enum LlmEmbedInput { + Text(String), + Batch(Vec), +} + +#[derive(Debug, Deserialize, Serialize)] +struct LlmModelParams { + model: String, +} + +#[derive(Debug, Deserialize, Serialize)] +#[serde(rename_all = "snake_case")] +struct LlmAdminLoadParams { + model: String, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + temperature: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_k: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + top_p: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + repeat_penalty: Option, + #[serde(skip_serializing_if = "Option::is_none")] + #[serde(default)] + max_tokens: Option, +} + +#[derive(Debug, Clone)] +struct ProjectRecord { + id: Uuid, + owner_id: i32, + name: String, + description: Option, + created_at: DateTime, + updated_at: DateTime, +} + +impl ProjectRecord { + fn to_value(&self) -> Value { + json!({ + "id": self.id, + "owner_id": self.owner_id, + "name": self.name.clone(), + "description": self.description.clone(), + "created_at": self.created_at.to_rfc3339(), + "updated_at": self.updated_at.to_rfc3339(), + }) + } +} + +fn normalize_project_name(name: &str) -> std::result::Result { + let trimmed = name.trim(); + if trimmed.is_empty() { + return Err(RpcMethodError::new( + -32602, + "project name is required", + None, + )); + } + if trimmed.len() > 128 { + return Err(RpcMethodError::new( + -32602, + "project name must be at most 128 characters", + Some(json!({ "max": 128 })), + )); + } + Ok(trimmed.to_string()) +} + +fn truncate_description(value: &str) -> String { + let trimmed = value.trim(); + let mut result = String::with_capacity(trimmed.len().min(512)); + for ch in trimmed.chars().take(512) { + result.push(ch); + } + result +} + +fn project_directory_relative(project_id: &Uuid) -> PathBuf { + PathBuf::from("projects").join(project_id.to_string()) +} + +fn parse_project_id(value: &str) -> std::result::Result { + Uuid::parse_str(value).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid project identifier", + Some(json!({ "detail": err.to_string() })), + ) + }) +} + +fn normalize_project_path(path: &str) -> std::result::Result { + let trimmed = path.trim(); + if trimmed.is_empty() { + return Err(RpcMethodError::new( + -32602, + "project path is required", + None, + )); + } + if trimmed.len() > 512 { + return Err(RpcMethodError::new( + -32602, + "project path must be at most 512 characters", + Some(json!({ "max": 512 })), + )); + } + let candidate = Path::new(trimmed); + if candidate.is_absolute() { + return Err(RpcMethodError::new( + -32602, + "project paths must be relative", + Some(json!({ "path": trimmed })), + )); + } + let mut normalized = PathBuf::new(); + for component in candidate.components() { + match component { + Component::Normal(part) => normalized.push(part), + Component::CurDir => continue, + _ => { + return Err(RpcMethodError::new( + -32602, + "project path cannot traverse parents", + Some(json!({ "path": trimmed })), + )) + } + } + } + if normalized.as_os_str().is_empty() { + return Err(RpcMethodError::new( + -32602, + "project path cannot resolve to empty", + Some(json!({ "path": trimmed })), + )); + } + Ok(normalized) +} + +async fn create_project( + pool: &PgPool, + ctx: &RequestContext, + name: &str, + description: Option<&str>, +) -> std::result::Result { + let row = sqlx::query( + "INSERT INTO projects (user_id, name, description) VALUES ($1, $2, $3) RETURNING id, user_id, name, description, created_at, updated_at", + ) + .bind(ctx.user_id) + .bind(name) + .bind(description) + .fetch_one(pool) + .await + .map_err(|err| match &err { + SqlxError::Database(db_err) if db_err.code().as_deref() == Some("23505") => { + RpcMethodError::new( + -32052, + "a project with this name already exists", + Some(json!({ "name": name })), + ) + } + _ => RpcMethodError::internal(&format!("failed to create project: {err}")), + })?; + + Ok(ProjectRecord { + id: row.get("id"), + owner_id: row.get("user_id"), + name: row.get("name"), + description: row.get("description"), + created_at: row.get("created_at"), + updated_at: row.get("updated_at"), + }) +} + +async fn list_projects( + pool: &PgPool, + ctx: &RequestContext, +) -> std::result::Result, RpcMethodError> { + let rows = if ctx.is_admin() { + sqlx::query( + "SELECT id, user_id, name, description, created_at, updated_at FROM projects ORDER BY created_at DESC", + ) + .fetch_all(pool) + .await + } else { + sqlx::query( + "SELECT id, user_id, name, description, created_at, updated_at FROM projects WHERE user_id = $1 ORDER BY created_at DESC", + ) + .bind(ctx.user_id) + .fetch_all(pool) + .await + } + .map_err(|err| RpcMethodError::internal(&format!("failed to list projects: {err}")))?; + + Ok(rows + .into_iter() + .map(|row| { + let created: DateTime = row.get("created_at"); + let updated: DateTime = row.get("updated_at"); + json!({ + "id": row.get::("id"), + "owner_id": row.get::("user_id"), + "name": row.get::("name"), + "description": row.get::, _>("description"), + "created_at": created.to_rfc3339(), + "updated_at": updated.to_rfc3339(), + }) + }) + .collect()) +} + +async fn load_project( + pool: &PgPool, + ctx: &RequestContext, + project_id: &Uuid, +) -> std::result::Result { + let row = sqlx::query( + "SELECT id, user_id, name, description, created_at, updated_at FROM projects WHERE id = $1", + ) + .bind(project_id) + .fetch_optional(pool) + .await + .map_err(|err| RpcMethodError::internal(&format!("failed to load project: {err}")))?; + + let row = row.ok_or_else(|| RpcMethodError::new(-32055, "project not found", None))?; + let owner_id: i32 = row.get("user_id"); + if owner_id != ctx.user_id && !ctx.is_admin() { + return Err(RpcMethodError::forbidden("project access denied")); + } + + Ok(ProjectRecord { + id: row.get("id"), + owner_id, + name: row.get("name"), + description: row.get("description"), + created_at: row.get("created_at"), + updated_at: row.get("updated_at"), + }) +} + +async fn project_files( + pool: &PgPool, + project_id: &Uuid, + include_content: bool, +) -> std::result::Result, RpcMethodError> { + let rows = sqlx::query( + "SELECT path, size, sha256, updated_at, content FROM project_files WHERE project_id = $1 ORDER BY path", + ) + .bind(project_id) + .fetch_all(pool) + .await + .map_err(|err| RpcMethodError::internal(&format!("failed to load project files: {err}")))?; + + let mut files = Vec::with_capacity(rows.len()); + for row in rows { + let path: String = row.get("path"); + let size: i64 = row.get("size"); + let sha: Vec = row.get("sha256"); + let updated: DateTime = row.get("updated_at"); + let mut object = serde_json::Map::new(); + object.insert("path".to_string(), Value::String(path)); + object.insert("size".to_string(), Value::Number(size.into())); + object.insert("sha256".to_string(), Value::String(hex_encode(sha))); + object.insert( + "updated_at".to_string(), + Value::String(updated.to_rfc3339()), + ); + if include_content { + let content: Vec = row.get("content"); + object.insert("data".to_string(), Value::String(BASE64.encode(content))); + } + files.push(Value::Object(object)); + } + Ok(files) +} + +async fn delete_project( + pool: &PgPool, + project_id: &Uuid, +) -> std::result::Result<(), RpcMethodError> { + sqlx::query("DELETE FROM projects WHERE id = $1") + .bind(project_id) + .execute(pool) + .await + .map_err(|err| RpcMethodError::internal(&format!("failed to delete project: {err}")))?; + Ok(()) +} + +async fn save_project_file( + pool: &PgPool, + project_id: &Uuid, + path: &Path, + data: &[u8], + sha256: &[u8], +) -> std::result::Result { + let path_str = path.to_string_lossy().to_string(); + let row = sqlx::query( + "INSERT INTO project_files (project_id, path, content, sha256, size) VALUES ($1, $2, $3, $4, $5) + ON CONFLICT (project_id, path) DO UPDATE SET content = EXCLUDED.content, sha256 = EXCLUDED.sha256, size = EXCLUDED.size, updated_at = NOW() + RETURNING updated_at", + ) + .bind(project_id) + .bind(&path_str) + .bind(data) + .bind(sha256) + .bind(data.len() as i64) + .fetch_one(pool) + .await + .map_err(|err| RpcMethodError::internal(&format!("failed to save project file: {err}")))?; + + let updated: DateTime = row.get("updated_at"); + Ok(json!({ + "status": "ok", + "path": path_str, + "size": data.len() as i64, + "sha256": hex_encode(sha256), + "updated_at": updated.to_rfc3339(), + })) +} + +async fn read_project_file( + pool: &PgPool, + project_id: &Uuid, + path: &Path, +) -> std::result::Result { + let path_str = path.to_string_lossy().to_string(); + let row = sqlx::query( + "SELECT content, size, sha256, updated_at FROM project_files WHERE project_id = $1 AND path = $2", + ) + .bind(project_id) + .bind(&path_str) + .fetch_optional(pool) + .await + .map_err(|err| RpcMethodError::internal(&format!("failed to read project file: {err}")))?; + + let row = row.ok_or_else(|| { + RpcMethodError::new( + -32052, + "project file not found", + Some(json!({ "path": path_str.clone() })), + ) + })?; + let content: Vec = row.get("content"); + let sha: Vec = row.get("sha256"); + let updated: DateTime = row.get("updated_at"); + let size: i64 = row.get("size"); + + Ok(json!({ + "path": path_str, + "data": BASE64.encode(content), + "size": size, + "sha256": hex_encode(sha), + "updated_at": updated.to_rfc3339(), + })) +} + +async fn delete_project_file( + pool: &PgPool, + project_id: &Uuid, + path: &Path, +) -> std::result::Result<(), RpcMethodError> { + let path_str = path.to_string_lossy().to_string(); + let result = sqlx::query("DELETE FROM project_files WHERE project_id = $1 AND path = $2") + .bind(project_id) + .bind(&path_str) + .execute(pool) + .await + .map_err(|err| { + RpcMethodError::internal(&format!("failed to delete project file: {err}")) + })?; + if result.rows_affected() == 0 { + return Err(RpcMethodError::new( + -32052, + "project file not found", + Some(json!({ "path": path_str })), + )); + } + Ok(()) +} + +async fn record_project_activity( + pool: &PgPool, + project_id: Uuid, + user_id: i32, + action: &str, + detail: Option, +) -> Result<(), SqlxError> { + sqlx::query( + "INSERT INTO project_activity (project_id, user_id, action, detail) VALUES ($1, $2, $3, $4)", + ) + .bind(project_id) + .bind(user_id) + .bind(action) + .bind(Json(detail.unwrap_or(Value::Null))) + .execute(pool) + .await + .map(|_| ()) +} + +fn map_db_activity_error(err: SqlxError, message: &str) -> RpcMethodError { + RpcMethodError::internal(&format!("{message}: {err}")) +} + +fn parse_params Deserialize<'a>>( + params: Option, +) -> std::result::Result { + let value = params.unwrap_or_else(|| Value::Object(Default::default())); + serde_json::from_value(value).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid params", + Some(json!({ "detail": err.to_string() })), + ) + }) +} + +fn enrich_agent_metadata(metadata: Option, ctx: &RequestContext) -> Option { + let mut map = metadata + .and_then(|value| value.as_object().cloned()) + .unwrap_or_default(); + map.insert( + "requested_by".to_string(), + Value::String(ctx.username.clone()), + ); + map.insert("requested_by_id".to_string(), json!(ctx.user_id)); + map.insert( + "auth_source".to_string(), + Value::String(ctx.auth_source().to_string()), + ); + map.insert( + "role".to_string(), + Value::String(ctx.role.as_str().to_string()), + ); + if let Some(api_key_id) = ctx.api_key_id { + map.insert("api_key_id".to_string(), json!(api_key_id)); + } + Some(Value::Object(map)) +} + +fn build_agent_context( + sandbox: &SandboxFs, + params: Option, +) -> std::result::Result { + let mut context = AgentContext::default(); + if let Some(ctx) = params { + context.notes = ctx.notes; + for file in ctx.files { + let (entry, note) = resolve_agent_context_file(sandbox, file)?; + if let Some(extra) = note { + context.notes.push(extra); + } + context.files.push(entry); + } + } + Ok(context) +} + +fn resolve_agent_context_file( + sandbox: &SandboxFs, + params: AgentDispatchContextFileParams, +) -> std::result::Result<(AgentContextFile, Option), SandboxError> { + let limit = params.max_bytes.unwrap_or(64 * 1024); + let title = params + .title + .clone() + .or_else(|| params.path.clone()) + .unwrap_or_else(|| "context".to_string()); + + if let Some(content_base64) = params.content_base64 { + let mut bytes = BASE64.decode(content_base64.as_bytes()).map_err(|err| { + SandboxError::InvalidOperation(format!("invalid base64 inline content: {err}")) + })?; + let mut note = None; + if bytes.len() > limit { + bytes.truncate(limit); + note = Some(format!( + "Inline content '{}' truncated to {} bytes", + title, limit + )); + } + let encoding = params.encoding.unwrap_or_else(|| "utf-8".to_string()); + let content = if encoding.eq_ignore_ascii_case("utf-8") { + match String::from_utf8(bytes) { + Ok(text) => AgentFileContent::Utf8(text), + Err(err) => { + let bytes = err.into_bytes(); + let encoded = BASE64.encode(&bytes); + let detail = format!( + "Inline content '{}' was not valid UTF-8; provided as base64", + title + ); + note = Some(match note { + Some(existing) => format!("{existing}; {detail}"), + None => detail, + }); + AgentFileContent::Base64(encoded) + } + } + } else { + AgentFileContent::Base64(BASE64.encode(&bytes)) + }; + return Ok(( + AgentContextFile { + path: params.path, + title, + content, + }, + note, + )); + } + + let path = params.path.ok_or_else(|| { + SandboxError::InvalidOperation( + "context file path is required when no inline content is provided".to_string(), + ) + })?; + let mut data = sandbox.read(Path::new(&path))?; + let mut note = None; + if data.len() > limit { + data.truncate(limit); + note = Some(format!( + "File '{}' truncated to {} bytes for agent context", + path, limit + )); + } + let encoding = params.encoding.unwrap_or_else(|| "utf-8".to_string()); + let content = if encoding.eq_ignore_ascii_case("base64") { + AgentFileContent::Base64(BASE64.encode(&data)) + } else { + match String::from_utf8(data) { + Ok(text) => AgentFileContent::Utf8(text), + Err(err) => { + let bytes = err.into_bytes(); + let encoded = BASE64.encode(&bytes); + let detail = format!( + "File '{}' contained non UTF-8 data; provided as base64", + path + ); + note = Some(match note { + Some(existing) => format!("{existing}; {detail}"), + None => detail, + }); + AgentFileContent::Base64(encoded) + } + } + }; + Ok(( + AgentContextFile { + path: Some(path), + title, + content, + }, + note, + )) +} + +#[derive(Debug, Deserialize)] +struct RpcRequest { + jsonrpc: String, + method: String, + params: Option, + id: Value, +} + +#[derive(Debug, Serialize)] +struct RpcResponse { + jsonrpc: &'static str, + #[serde(skip_serializing_if = "Option::is_none")] + result: Option, + #[serde(skip_serializing_if = "Option::is_none")] + error: Option, + id: Value, +} + +impl RpcResponse { + fn success(id: Value, result: Value) -> Self { + Self { + jsonrpc: "2.0", + result: Some(result), + error: None, + id, + } + } + + fn error(id: Value, code: i64, message: &str, data: Option) -> Self { + Self { + jsonrpc: "2.0", + result: None, + error: Some(RpcError { + code, + message: message.to_string(), + data, + }), + id, + } + } +} + +#[derive(Debug, Serialize)] +struct RpcError { + code: i64, + message: String, + #[serde(skip_serializing_if = "Option::is_none")] + data: Option, +} + +#[derive(Debug)] +struct RpcMethodError { + code: i64, + message: String, + data: Option, +} + +impl RpcMethodError { + fn new(code: i64, message: &str, data: Option) -> Self { + Self { + code, + message: message.to_string(), + data, + } + } + + fn from_sandbox(code: i64, message: &str, err: sandbox::SandboxError) -> Self { + Self { + code, + message: message.to_string(), + data: Some(json!({ "detail": err.to_string() })), + } + } + + fn unauthorized(message: &str) -> Self { + Self::new(-32090, message, None) + } + + fn forbidden(message: &str) -> Self { + Self::new(-32091, message, None) + } + + fn internal(detail: &str) -> Self { + Self::new(-32603, "internal error", Some(json!({ "detail": detail }))) + } +} + +#[derive(Debug, Deserialize)] +struct FsPathParams { + path: String, +} + +#[derive(Debug, Deserialize)] +struct FsWriteParams { + path: String, + data: String, +} + +#[derive(Debug, Deserialize)] +struct ProjectCreateParams { + name: String, + #[serde(default)] + description: Option, +} + +#[derive(Debug, Deserialize)] +struct ProjectIdParams { + project_id: String, +} + +#[derive(Debug, Deserialize)] +struct ProjectOpenParams { + project_id: String, + #[serde(default)] + include_content: Option, +} + +#[derive(Debug, Deserialize)] +struct ProjectFileSaveParams { + project_id: String, + path: String, + data: String, + #[serde(default)] + encoding: Option, + #[serde(default)] + message: Option, +} + +#[derive(Debug, Deserialize)] +struct ProjectFilePathParams { + project_id: String, + path: String, +} + +#[derive(Debug, Deserialize)] +struct RunExecParams { + program: String, + #[serde(default)] + args: Vec, + #[serde(default)] + env: Vec, + #[serde(default)] + stdin: Option, + #[serde(default)] + cwd: Option, + #[serde(default)] + timeout_ms: Option, +} + +impl RunExecParams { + fn into_request(self) -> std::result::Result { + let mut request = RunRequest::new(self.program); + if !self.args.is_empty() { + request.args = self.args; + } + if !self.env.is_empty() { + request.env = self + .env + .into_iter() + .map(|pair| (pair.key, pair.value)) + .collect(); + } + if let Some(stdin) = self.stdin { + if !stdin.is_empty() { + let data = BASE64.decode(stdin.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + request.stdin = Some(data); + } + } + if let Some(cwd) = self.cwd { + if !cwd.is_empty() { + request.working_dir = Some(cwd); + } + } + if let Some(timeout_ms) = self.timeout_ms { + request.timeout = Some(Duration::from_millis(timeout_ms)); + } + Ok(request) + } +} + +#[derive(Debug, Deserialize, Clone)] +struct RunEnvVar { + key: String, + value: String, +} + +#[derive(Debug, Deserialize)] +struct MicroStartParams { + image: String, + #[serde(default)] + init_script: Option, +} + +#[derive(Debug, Deserialize)] +struct MicroExecuteParams { + vm_id: String, + code: String, + #[serde(default)] + timeout_ms: Option, +} + +#[derive(Debug, Deserialize)] +struct MicroStopParams { + vm_id: String, +} + +#[derive(Debug, Deserialize)] +struct RawMicroImage { + name: String, + command: String, + #[serde(default)] + args: Vec, + #[serde(default)] + extension: Option, + #[serde(default)] + env: Vec, +} + +#[derive(Debug, Deserialize)] +struct AgentDispatchParams { + agent: AgentKind, + objective: String, + #[serde(default)] + context: Option, + #[serde(default)] + model: Option, + #[serde(default)] + metadata: Option, + #[serde(default)] + parameters: Option, +} + +#[derive(Debug, Deserialize, Default)] +struct AgentDispatchContextParams { + #[serde(default)] + notes: Vec, + #[serde(default)] + files: Vec, +} + +#[derive(Debug, Deserialize)] +struct AgentDispatchContextFileParams { + #[serde(default)] + path: Option, + #[serde(default)] + title: Option, + #[serde(default)] + encoding: Option, + #[serde(default)] + max_bytes: Option, + #[serde(default)] + content_base64: Option, +} + +#[derive(Debug, Deserialize)] +struct AgentParameterOverrides { + #[serde(default)] + temperature: Option, + #[serde(default)] + max_tokens: Option, + #[serde(default)] + top_p: Option, +} + +impl AgentParameterOverrides { + fn into_parameters(self) -> AgentParameters { + let mut params = AgentParameters::default(); + if let Some(temp) = self.temperature { + params.temperature = temp; + } + if let Some(max_tokens) = self.max_tokens { + params.max_tokens = Some(max_tokens); + } + if let Some(top_p) = self.top_p { + params.top_p = top_p; + } + params + } +} + +#[derive(Debug, Deserialize)] +struct AgentStatusParams { + task_id: String, +} + +#[derive(Debug, Deserialize)] +struct AgentHistoryParams { + #[serde(default)] + limit: Option, +} + +#[derive(Debug, Deserialize)] +struct WasmInvokeParams { + #[serde(default)] + module_path: Option, + #[serde(default)] + module_bytes: Option, + function: String, + #[serde(default)] + params: Vec, + #[serde(default)] + fuel: Option, + #[serde(default)] + memory_limit: Option, + #[serde(default)] + table_elements_limit: Option, +} + +#[derive(Debug, Deserialize)] +#[serde(tag = "type", content = "value")] +enum WasmParam { + #[serde(rename = "i32")] + I32(i32), + #[serde(rename = "i64")] + I64(i64), + #[serde(rename = "f32")] + F32(f32), + #[serde(rename = "f64")] + F64(f64), +} + +impl WasmParam { + fn into_value(self) -> std::result::Result { + Ok(match self { + WasmParam::I32(value) => WasmValue::I32(value), + WasmParam::I64(value) => WasmValue::I64(value), + WasmParam::F32(value) => WasmValue::F32(value), + WasmParam::F64(value) => WasmValue::F64(value), + }) + } +} + +fn resolve_wasm_module( + params: &WasmInvokeParams, +) -> std::result::Result { + match (¶ms.module_path, ¶ms.module_bytes) { + (Some(_), Some(_)) => Err(RpcMethodError::new( + -32602, + "specify either module_path or module_bytes", + None, + )), + (None, None) => Err(RpcMethodError::new( + -32602, + "missing wasm module source", + None, + )), + (Some(path), None) => Ok(WasmModuleSource::from_path(path.clone())), + (None, Some(bytes)) => { + if bytes.is_empty() { + return Err(RpcMethodError::new( + -32602, + "module_bytes must not be empty", + None, + )); + } + let decoded = BASE64.decode(bytes.as_bytes()).map_err(|err| { + RpcMethodError::new( + -32602, + "invalid base64 payload", + Some(json!({ "detail": err.to_string() })), + ) + })?; + Ok(WasmModuleSource::from_bytes(decoded)) + } + } +} + +fn wasm_value_to_json(value: WasmValue) -> Value { + match value { + WasmValue::I32(v) => json!({ "type": "i32", "value": v }), + WasmValue::I64(v) => json!({ "type": "i64", "value": v }), + WasmValue::F32(v) => json!({ "type": "f32", "value": v }), + WasmValue::F64(v) => json!({ "type": "f64", "value": v }), + } +} + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn normalize_project_name_trims_and_limits_length() { + assert_eq!(normalize_project_name(" demo ").unwrap(), "demo"); + assert!(normalize_project_name("").is_err()); + let oversized = "a".repeat(129); + assert!(normalize_project_name(&oversized).is_err()); + } + + #[test] + fn normalize_project_path_rejects_parent_traversal() { + assert!(normalize_project_path("../secret").is_err()); + assert!(normalize_project_path("/absolute").is_err()); + let path = normalize_project_path("src/lib.rs").expect("valid path"); + assert_eq!(path.to_string_lossy(), "src/lib.rs"); + } +} diff --git a/apps/auth/Cargo.toml b/apps/auth/Cargo.toml new file mode 100644 index 0000000..adbca58 --- /dev/null +++ b/apps/auth/Cargo.toml @@ -0,0 +1,22 @@ +[package] +name = "auth" +version = "0.1.0" +edition = "2021" + +[dependencies] +anyhow = { workspace = true } +axum = { workspace = true } +bcrypt = { workspace = true } +chrono = { workspace = true } +jsonwebtoken = { workspace = true } +hex = { workspace = true } +rand = { workspace = true } +serde = { workspace = true } +serde_json = { workspace = true } +sha2 = { workspace = true } +sqlx = { workspace = true } +tokio = { workspace = true } +tracing = { workspace = true } +tracing-subscriber = { workspace = true } +uuid = { workspace = true } +thiserror = { workspace = true } diff --git a/apps/auth/src/main.rs b/apps/auth/src/main.rs new file mode 100644 index 0000000..34d89bc --- /dev/null +++ b/apps/auth/src/main.rs @@ -0,0 +1,464 @@ +use std::net::SocketAddr; +use std::sync::Arc; + +use axum::extract::{Path, State}; +use axum::http::{HeaderMap, StatusCode}; +use axum::response::{IntoResponse, Response}; +use axum::routing::{delete, get, post}; +use axum::{Json, Router}; +use chrono::{Duration, Utc}; +use jsonwebtoken::{decode, encode, Algorithm, DecodingKey, EncodingKey, Header, Validation}; +use serde::{Deserialize, Serialize}; +use sqlx::postgres::PgPoolOptions; +use sqlx::{PgPool, Row}; +use tower_http::trace::TraceLayer; +use tracing::{dispatcher, error, info}; +use uuid::Uuid; + +use hex::encode as hex_encode; +use rand::rngs::OsRng; +use rand::RngCore; +use sha2::{Digest, Sha256}; + +#[derive(Clone)] +struct AppState { + pool: PgPool, + jwt: JwtConfig, +} + +#[derive(Clone)] +struct JwtConfig { + secret: Arc<[u8]>, + expiration: Duration, + issuer: String, +} + +impl JwtConfig { + fn from_env() -> anyhow::Result { + let secret = std::env::var("AUTH_JWT_SECRET") + .map_err(|_| anyhow::anyhow!("AUTH_JWT_SECRET environment variable is required"))?; + let expiration_minutes = std::env::var("AUTH_JWT_EXP_MINUTES") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(60); + let issuer = + std::env::var("AUTH_JWT_ISSUER").unwrap_or_else(|_| "cyber-dev-studio".to_string()); + Ok(Self { + secret: Arc::from(secret.into_bytes()), + expiration: Duration::minutes(expiration_minutes), + issuer, + }) + } + + fn validation(&self) -> Validation { + let mut validation = Validation::new(Algorithm::HS256); + validation + .set_required_spec_claims(&["exp", "iat", "sub", "iss"]) + .expect("required claim configuration"); + validation.iss = Some(self.issuer.clone()); + validation + } +} + +#[derive(Debug, Serialize, Deserialize)] +struct Claims { + sub: i32, + username: String, + role: String, + exp: usize, + iat: usize, + iss: String, + jti: String, +} + +#[derive(Debug)] +struct AuthenticatedUser { + user_id: i32, + username: String, + role: String, +} + +#[tokio::main] +async fn main() -> anyhow::Result<()> { + init_tracing(); + let bind_addr = resolve_bind_address()?; + let pool = build_pool().await?; + let jwt = JwtConfig::from_env()?; + + let state = AppState { pool, jwt }; + + let app = Router::new() + .route("/health", get(health)) + .route("/auth/register", post(register_user)) + .route("/auth/login", post(login_user)) + .route("/auth/api-keys", get(list_api_keys).post(create_api_key)) + .route("/auth/api-keys/:id", delete(delete_api_key)) + .with_state(state) + .layer(TraceLayer::new_for_http()); + + info!("binding", %bind_addr, "auth service starting"); + axum::Server::bind(&bind_addr) + .serve(app.into_make_service()) + .await?; + Ok(()) +} + +fn init_tracing() { + if dispatcher::has_been_set() { + return; + } + let subscriber = tracing_subscriber::fmt() + .with_env_filter( + tracing_subscriber::EnvFilter::try_from_default_env() + .unwrap_or_else(|_| "info,tower_http=info".into()), + ) + .json() + .finish(); + if let Err(err) = tracing::subscriber::set_global_default(subscriber) { + eprintln!("failed to install tracing subscriber: {err}"); + } +} + +fn resolve_bind_address() -> anyhow::Result { + let raw = std::env::var("AUTH_BIND_ADDR").unwrap_or_else(|_| "0.0.0.0:6971".to_string()); + Ok(raw.parse()?) +} + +async fn build_pool() -> anyhow::Result { + let database_url = std::env::var("DATABASE_URL") + .map_err(|_| anyhow::anyhow!("DATABASE_URL environment variable is required"))?; + let max_connections = std::env::var("DATABASE_MAX_CONNECTIONS") + .ok() + .and_then(|v| v.parse::().ok()) + .unwrap_or(5); + let pool = PgPoolOptions::new() + .max_connections(max_connections) + .acquire_timeout(std::time::Duration::from_secs(10)) + .connect(&database_url) + .await?; + Ok(pool) +} + +async fn health() -> impl IntoResponse { + (StatusCode::OK, Json(serde_json::json!({ "status": "ok" }))) +} + +async fn register_user( + State(state): State, + Json(payload): Json, +) -> Result, AuthError> { + validate_role(&payload.role)?; + if payload.password.len() < 12 { + return Err(AuthError::BadRequest( + "password must contain at least 12 characters".to_string(), + )); + } + + let hashed = bcrypt::hash(&payload.password, bcrypt::DEFAULT_COST) + .map_err(|err| AuthError::Internal(err.to_string()))?; + let role = payload.role.unwrap_or_else(|| "developer".to_string()); + + let rec = sqlx::query( + "INSERT INTO users (username, password_hash, role, token_balance) VALUES ($1, $2, $3, $4) RETURNING id", + ) + .bind(&payload.username) + .bind(&hashed) + .bind(&role) + .bind(payload.initial_tokens.unwrap_or(0_i64)) + .fetch_one(&state.pool) + .await + .map_err(|err| match err { + sqlx::Error::Database(db_err) if db_err.code().as_deref() == Some("23505") => { + AuthError::Conflict(format!("user '{}' already exists", payload.username)) + } + other => AuthError::Internal(other.to_string()), + })?; + + let id: i32 = rec.get("id"); + Ok(Json(RegisterResponse { user_id: id })) +} + +async fn login_user( + State(state): State, + Json(payload): Json, +) -> Result, AuthError> { + let row = sqlx::query("SELECT id, password_hash, role FROM users WHERE username = $1") + .bind(&payload.username) + .fetch_one(&state.pool) + .await + .map_err(|err| match err { + sqlx::Error::RowNotFound => AuthError::Unauthorized("invalid credentials".to_string()), + other => AuthError::Internal(other.to_string()), + })?; + + let stored_hash: String = row.get("password_hash"); + if !bcrypt::verify(&payload.password, &stored_hash) + .map_err(|err| AuthError::Internal(err.to_string()))? + { + return Err(AuthError::Unauthorized("invalid credentials".to_string())); + } + let user_id: i32 = row.get("id"); + let role: String = row.get("role"); + + let claims = Claims::new(user_id, &payload.username, &role, &state.jwt); + let token = encode( + &Header::default(), + &claims, + &EncodingKey::from_secret(&state.jwt.secret), + ) + .map_err(|err| AuthError::Internal(err.to_string()))?; + + Ok(Json(LoginResponse { + token, + expires_at: chrono::DateTime::::from_timestamp(claims.exp as i64, 0) + .expect("valid expiration timestamp"), + })) +} + +async fn list_api_keys( + State(state): State, + headers: HeaderMap, +) -> Result, AuthError> { + let user = authenticate(&headers, &state).await?; + let records = sqlx::query( + "SELECT id, name, created_at, last_used_at FROM api_keys WHERE user_id = $1 ORDER BY created_at DESC", + ) + .bind(user.user_id) + .fetch_all(&state.pool) + .await + .map_err(|err| AuthError::Internal(err.to_string()))?; + + let keys = records + .into_iter() + .map(|row| ApiKeySummary { + id: row.get("id"), + name: row.get("name"), + created_at: row.get("created_at"), + last_used_at: row.get("last_used_at"), + }) + .collect(); + + Ok(Json(ListApiKeysResponse { keys })) +} + +async fn create_api_key( + State(state): State, + headers: HeaderMap, + Json(payload): Json, +) -> Result, AuthError> { + let user = authenticate(&headers, &state).await?; + let mut name = payload + .name + .unwrap_or_else(|| format!("key-{}", Utc::now().timestamp())); + name.truncate(128); + let trimmed = name.trim(); + if trimmed.is_empty() { + return Err(AuthError::BadRequest("name must not be empty".to_string())); + } + let normalized_name = trimmed.to_string(); + + let api_key = generate_api_key(); + let hash = hash_api_key(&api_key); + + let record = sqlx::query( + "INSERT INTO api_keys (user_id, name, api_key_hash) VALUES ($1, $2, $3) RETURNING id, created_at", + ) + .bind(user.user_id) + .bind(&normalized_name) + .bind(&hash) + .fetch_one(&state.pool) + .await + .map_err(|err| AuthError::Internal(err.to_string()))?; + + Ok(Json(CreateApiKeyResponse { + id: record.get("id"), + name: normalized_name, + key: api_key, + created_at: record.get("created_at"), + })) +} + +async fn delete_api_key( + State(state): State, + headers: HeaderMap, + Path(id): Path, +) -> Result { + let user = authenticate(&headers, &state).await?; + let result = sqlx::query("DELETE FROM api_keys WHERE id = $1 AND user_id = $2") + .bind(id) + .bind(user.user_id) + .execute(&state.pool) + .await + .map_err(|err| AuthError::Internal(err.to_string()))?; + + if result.rows_affected() == 0 { + return Err(AuthError::NotFound("api key not found".to_string())); + } + + Ok(StatusCode::NO_CONTENT) +} + +async fn authenticate( + headers: &HeaderMap, + state: &AppState, +) -> Result { + let authorization = headers + .get(axum::http::header::AUTHORIZATION) + .ok_or_else(|| AuthError::Unauthorized("missing authorization header".to_string()))?; + let authorization = authorization + .to_str() + .map_err(|_| AuthError::Unauthorized("invalid authorization header".to_string()))?; + let token = authorization + .strip_prefix("Bearer ") + .ok_or_else(|| AuthError::Unauthorized("unsupported authorization scheme".to_string()))?; + + let validation = state.jwt.validation(); + let token_data = decode::( + token, + &DecodingKey::from_secret(&state.jwt.secret), + &validation, + ) + .map_err(|_| AuthError::Unauthorized("invalid token".to_string()))?; + let claims = token_data.claims; + + let row = sqlx::query("SELECT username, role FROM users WHERE id = $1") + .bind(claims.sub) + .fetch_one(&state.pool) + .await + .map_err(|err| match err { + sqlx::Error::RowNotFound => AuthError::Unauthorized("user not found".to_string()), + other => AuthError::Internal(other.to_string()), + })?; + + Ok(AuthenticatedUser { + user_id: claims.sub, + username: row.get("username"), + role: row.get("role"), + }) +} + +fn generate_api_key() -> String { + let mut bytes = [0u8; 32]; + OsRng.fill_bytes(&mut bytes); + format!("cds_{}", hex_encode(bytes)) +} + +fn hash_api_key(key: &str) -> String { + let mut hasher = Sha256::new(); + hasher.update(key.as_bytes()); + hex_encode(hasher.finalize()) +} + +impl Claims { + fn new(user_id: i32, username: &str, role: &str, jwt: &JwtConfig) -> Self { + let now = Utc::now(); + let exp = now + jwt.expiration; + Self { + sub: user_id, + username: username.to_string(), + role: role.to_string(), + exp: exp.timestamp() as usize, + iat: now.timestamp() as usize, + iss: jwt.issuer.clone(), + jti: Uuid::new_v4().to_string(), + } + } +} + +fn validate_role(role: &Option) -> Result<(), AuthError> { + if let Some(role) = role { + match role.as_str() { + "admin" | "developer" | "viewer" => Ok(()), + _ => Err(AuthError::BadRequest(format!( + "unsupported role '{}'", + role + ))), + } + } else { + Ok(()) + } +} + +#[derive(Debug, Deserialize)] +struct RegisterRequest { + username: String, + password: String, + role: Option, + initial_tokens: Option, +} + +#[derive(Debug, Serialize)] +struct RegisterResponse { + user_id: i32, +} + +#[derive(Debug, Deserialize)] +struct LoginRequest { + username: String, + password: String, +} + +#[derive(Debug, Serialize)] +struct LoginResponse { + token: String, + expires_at: chrono::DateTime, +} + +#[derive(Debug, Deserialize)] +struct CreateApiKeyRequest { + #[serde(default)] + name: Option, +} + +#[derive(Debug, Serialize)] +struct CreateApiKeyResponse { + id: Uuid, + name: String, + key: String, + created_at: chrono::DateTime, +} + +#[derive(Debug, Serialize)] +struct ListApiKeysResponse { + keys: Vec, +} + +#[derive(Debug, Serialize)] +struct ApiKeySummary { + id: Uuid, + name: String, + created_at: chrono::DateTime, + #[serde(skip_serializing_if = "Option::is_none")] + last_used_at: Option>, +} + +#[derive(Debug, thiserror::Error)] +enum AuthError { + #[error("bad request: {0}")] + BadRequest(String), + #[error("unauthorized: {0}")] + Unauthorized(String), + #[error("conflict: {0}")] + Conflict(String), + #[error("not found: {0}")] + NotFound(String), + #[error("internal error: {0}")] + Internal(String), +} + +impl IntoResponse for AuthError { + fn into_response(self) -> Response { + let (status, message) = match &self { + AuthError::BadRequest(msg) => (StatusCode::BAD_REQUEST, msg.clone()), + AuthError::Unauthorized(msg) => (StatusCode::UNAUTHORIZED, msg.clone()), + AuthError::Conflict(msg) => (StatusCode::CONFLICT, msg.clone()), + AuthError::NotFound(msg) => (StatusCode::NOT_FOUND, msg.clone()), + AuthError::Internal(msg) => (StatusCode::INTERNAL_SERVER_ERROR, msg.clone()), + }; + error!("auth error", %message, kind = ?self); + let body = Json(serde_json::json!({ + "error": message, + })); + (status, body).into_response() + } +} diff --git a/apps/llmserver/package.json b/apps/llmserver/package.json new file mode 100644 index 0000000..2b7ea0f --- /dev/null +++ b/apps/llmserver/package.json @@ -0,0 +1,34 @@ +{ + "name": "cyberdevstudio-llmserver", + "version": "1.0.0", + "description": "CyberDevStudio LLM service wrapping node-llama-cpp with admin and accounting", + "main": "dist/server.js", + "scripts": { + "build": "tsc -p tsconfig.json", + "start": "node dist/server.js", + "dev": "ts-node --project tsconfig.json src/server.ts", + "lint": "tsc -p tsconfig.json --noEmit" + }, + "dependencies": { + "@huggingface/hub": "^0.15.0", + "cors": "^2.8.5", + "dotenv": "^16.4.5", + "express": "^4.19.2", + "express-ws": "^5.0.2", + "node-llama-cpp": "^2.8.0", + "pg": "^8.11.5", + "prom-client": "^15.1.1", + "uuid": "^9.0.1", + "ws": "^8.17.0", + "zod": "^3.22.4", + "jsonwebtoken": "^9.0.2" + }, + "devDependencies": { + "@types/express": "^4.17.21", + "@types/node": "^20.11.30", + "@types/ws": "^8.5.10", + "ts-node": "^10.9.2", + "typescript": "^5.4.3", + "@types/jsonwebtoken": "^9.0.6" + } +} diff --git a/apps/llmserver/src/auth.ts b/apps/llmserver/src/auth.ts new file mode 100644 index 0000000..6f6a33a --- /dev/null +++ b/apps/llmserver/src/auth.ts @@ -0,0 +1,21 @@ +import jwt from "jsonwebtoken"; +import { ServerConfig } from "./config"; + +export interface AdminClaims { + sub: string; + role: string; + exp: number; +} + +export function verifyAdminToken(config: ServerConfig, token: string): AdminClaims { + try { + const payload = jwt.verify(token, config.adminJwtSecret); + const claims = payload as AdminClaims; + if (claims.role !== "admin") { + throw new Error("Token does not grant admin privileges"); + } + return claims; + } catch (error) { + throw new Error(`Invalid admin token: ${String(error)}`); + } +} diff --git a/apps/llmserver/src/catalog.ts b/apps/llmserver/src/catalog.ts new file mode 100644 index 0000000..7157f61 --- /dev/null +++ b/apps/llmserver/src/catalog.ts @@ -0,0 +1,47 @@ +export interface ModelMetadata { + readonly name: string; + readonly displayName: string; + readonly huggingFaceRepo: string; + readonly file: string; + readonly context: number; + readonly costPerToken: number; +} + +export const MODEL_CATALOG: readonly ModelMetadata[] = [ + { + name: "deepseek-coder-1.3b", + displayName: "DeepSeek Coder 1.3B Instruct", + huggingFaceRepo: "deepseek-ai/deepseek-coder-1.3b-instruct-GGUF", + file: "deepseek-coder-1.3b-instruct.Q4_K_M.gguf", + context: 4096, + costPerToken: 0.00045 + }, + { + name: "nous-hermes-2-3b.Q4", + displayName: "Nous Hermes 2 3B Q4", + huggingFaceRepo: "TheBloke/Nous-Hermes-2-3B-GGUF", + file: "nous-hermes-llama2-3b.Q4_K_M.gguf", + context: 4096, + costPerToken: 0.00055 + }, + { + name: "bge-small-en-v1.5", + displayName: "BGE Small English v1.5", + huggingFaceRepo: "BAAI/bge-small-en-v1.5", + file: "bge-small-en-v1.5-q4_0.gguf", + context: 2048, + costPerToken: 0.00035 + }, + { + name: "tinyllama-1.1b-func", + displayName: "TinyLlama 1.1B Function", + huggingFaceRepo: "cognitivecomputations/TinyLlama-1.1B-Function-Call-GGUF", + file: "tinyllama-1.1b-chat-v1.0.Q4_K_M.gguf", + context: 2048, + costPerToken: 0.00030 + } +]; + +export function getModelMetadata(name: string): ModelMetadata | undefined { + return MODEL_CATALOG.find((model) => model.name === name); +} diff --git a/apps/llmserver/src/config.ts b/apps/llmserver/src/config.ts new file mode 100644 index 0000000..1284fa4 --- /dev/null +++ b/apps/llmserver/src/config.ts @@ -0,0 +1,63 @@ +import path from "node:path"; +import fs from "node:fs"; + +export interface ServerConfig { + readonly port: number; + readonly host: string; + readonly modelsDir: string; + readonly downloadsDir: string; + readonly llmThreadCount: number; + readonly llmBatchSize: number; + readonly databaseUrl: string; + readonly maxStreamingSeconds: number; + readonly adminJwtSecret: string; +} + +function ensureDirectory(dir: string): void { + if (!fs.existsSync(dir)) { + fs.mkdirSync(dir, { recursive: true }); + } +} + +export function loadConfig(): ServerConfig { + const port = parseInt(process.env.LLM_PORT ?? "6988", 10); + const host = process.env.LLM_HOST ?? "0.0.0.0"; + const modelsDir = process.env.LLM_MODELS_DIR ?? path.resolve(process.cwd(), "models"); + const downloadsDir = process.env.LLM_DOWNLOAD_DIR ?? path.join(modelsDir, "downloads"); + const llmThreadCount = parseInt(process.env.LLM_THREADS ?? "8", 10); + const llmBatchSize = parseInt(process.env.LLM_BATCH_SIZE ?? "1024", 10); + const maxStreamingSeconds = parseInt(process.env.LLM_MAX_STREAM_SECONDS ?? "30", 10); + const databaseUrl = process.env.DATABASE_URL; + const adminJwtSecret = process.env.LLM_ADMIN_JWT_SECRET ?? ""; + + if (!databaseUrl) { + throw new Error("DATABASE_URL environment variable is required for LLM server"); + } + if (!Number.isFinite(port) || port <= 0) { + throw new Error("LLM_PORT must be a positive integer"); + } + if (llmThreadCount <= 0) { + throw new Error("LLM_THREADS must be positive"); + } + if (llmBatchSize <= 0) { + throw new Error("LLM_BATCH_SIZE must be positive"); + } + if (!adminJwtSecret) { + throw new Error("LLM_ADMIN_JWT_SECRET must be configured"); + } + + ensureDirectory(modelsDir); + ensureDirectory(downloadsDir); + + return { + port, + host, + modelsDir, + downloadsDir, + llmThreadCount, + llmBatchSize, + databaseUrl, + maxStreamingSeconds, + adminJwtSecret + }; +} diff --git a/apps/llmserver/src/downloader.ts b/apps/llmserver/src/downloader.ts new file mode 100644 index 0000000..2ad0bf8 --- /dev/null +++ b/apps/llmserver/src/downloader.ts @@ -0,0 +1,47 @@ +import fs from "node:fs"; +import path from "node:path"; +import { downloadFile } from "@huggingface/hub"; +import { MODEL_CATALOG, getModelMetadata } from "./catalog"; +import { ServerConfig } from "./config"; + +export interface DownloadProgress { + model: string; + file: string; + size: number; +} + +export async function downloadModel(config: ServerConfig, modelName: string): Promise { + const metadata = getModelMetadata(modelName); + if (!metadata) { + throw new Error(`Model '${modelName}' is not supported`); + } + const targetDir = path.join(config.modelsDir, metadata.name); + const targetPath = path.join(targetDir, metadata.file); + if (!fs.existsSync(targetDir)) { + fs.mkdirSync(targetDir, { recursive: true }); + } + if (fs.existsSync(targetPath)) { + const stats = fs.statSync(targetPath); + return { model: metadata.name, file: targetPath, size: stats.size }; + } + + const tempPath = `${targetPath}.download`; + const download = await downloadFile({ + repo: metadata.huggingFaceRepo, + file: metadata.file, + localFile: tempPath + }); + if (!download) { + throw new Error(`Failed to download ${metadata.file} from HuggingFace repo ${metadata.huggingFaceRepo}`); + } + fs.renameSync(tempPath, targetPath); + const stats = fs.statSync(targetPath); + return { model: metadata.name, file: targetPath, size: stats.size }; +} + +export function listAvailableDownloads(): DownloadProgress[] { + return MODEL_CATALOG.map((metadata) => { + const file = `${metadata.huggingFaceRepo}/${metadata.file}`; + return { model: metadata.name, file, size: 0 }; + }); +} diff --git a/apps/llmserver/src/metrics.ts b/apps/llmserver/src/metrics.ts new file mode 100644 index 0000000..f50b6d3 --- /dev/null +++ b/apps/llmserver/src/metrics.ts @@ -0,0 +1,38 @@ +import client from "prom-client"; + +const register = new client.Registry(); + +export const requestCounter = new client.Counter({ + name: "llm_requests_total", + help: "Total number of LLM requests processed", + labelNames: ["endpoint", "model"] +}); + +export const tokenCounter = new client.Counter({ + name: "llm_tokens_generated", + help: "Total tokens generated across prompts and completions", + labelNames: ["model", "type"] +}); + +export const inferenceHistogram = new client.Histogram({ + name: "llm_inference_duration_seconds", + help: "Duration of inference operations", + labelNames: ["endpoint", "model"], + buckets: [0.1, 0.25, 0.5, 1, 2, 3, 5, 8, 13] +}); + +export const activeSessions = new client.Gauge({ + name: "llm_active_sessions", + help: "Active chat sessions" +}); + +register.registerMetric(requestCounter); +register.registerMetric(tokenCounter); +register.registerMetric(inferenceHistogram); +register.registerMetric(activeSessions); + +client.collectDefaultMetrics({ register }); + +export function metricsHandler(): Promise { + return register.metrics(); +} diff --git a/apps/llmserver/src/modelManager.ts b/apps/llmserver/src/modelManager.ts new file mode 100644 index 0000000..38d308a --- /dev/null +++ b/apps/llmserver/src/modelManager.ts @@ -0,0 +1,224 @@ +import fs from "node:fs"; +import path from "node:path"; +import EventEmitter from "node:events"; +import type { LlamaModelOptions } from "node-llama-cpp"; +import { MODEL_CATALOG, ModelMetadata, getModelMetadata } from "./catalog"; +import { ServerConfig } from "./config"; + +// node-llama-cpp does not ship TypeScript definitions for runtime types. +// eslint-disable-next-line @typescript-eslint/no-var-requires +const llama: any = require("node-llama-cpp"); + +type ChatSession = any; + +type LoadOverrides = { + temperature?: number; + topK?: number; + topP?: number; + repeatPenalty?: number; + maxTokens?: number; +}; + +export interface CompletionResult { + text: string; + promptTokens: number; + completionTokens: number; +} + +export interface EmbeddingResult { + embedding: number[]; + tokens: number; +} + +export interface ManagedModel { + metadata: ModelMetadata; + modelPath: string; + threads: number; + batchSize: number; + options: Required; + chat: ChatSession; + context: any; + model: any; + loadedAt: Date; +} + +export interface ModelStatus { + name: string; + displayName: string; + loaded: boolean; + contextSize: number; + loadedAt?: string; + threads?: number; + batchSize?: number; +} + +export class ModelManager extends EventEmitter { + private readonly config: ServerConfig; + private readonly models = new Map(); + + constructor(config: ServerConfig) { + super(); + this.config = config; + } + + listModels(): ModelStatus[] { + return MODEL_CATALOG.map((metadata) => { + const loaded = this.models.get(metadata.name); + return { + name: metadata.name, + displayName: metadata.displayName, + loaded: Boolean(loaded), + contextSize: metadata.context, + loadedAt: loaded?.loadedAt.toISOString(), + threads: loaded?.threads, + batchSize: loaded?.batchSize + }; + }); + } + + async loadModel(name: string, overrides?: LoadOverrides): Promise { + const metadata = requireMetadata(name); + const modelPath = this.resolveModelPath(metadata); + if (!fs.existsSync(modelPath)) { + throw new Error(`Model file not found at ${modelPath}. Download model before loading.`); + } + + const threads = this.config.llmThreadCount; + const batchSize = this.config.llmBatchSize; + const options: Required = { + temperature: overrides?.temperature ?? 0.2, + topK: overrides?.topK ?? 40, + topP: overrides?.topP ?? 0.9, + repeatPenalty: overrides?.repeatPenalty ?? 1.1, + maxTokens: overrides?.maxTokens ?? 1024 + }; + + const model = new llama.LlamaModel({ + modelPath, + contextSize: metadata.context, + gpuLayers: 0, + seed: Date.now() + } as LlamaModelOptions); + const context = new llama.LlamaContext({ model, threads, batchSize }); + const chat = new llama.LlamaChatSession({ context }); + + const managed: ManagedModel = { + metadata, + modelPath, + threads, + batchSize, + options, + chat, + context, + model, + loadedAt: new Date() + }; + this.models.set(metadata.name, managed); + this.emit("loaded", metadata.name); + return this.describe(metadata.name); + } + + describe(name: string): ModelStatus { + const metadata = requireMetadata(name); + const loaded = this.models.get(name); + return { + name: metadata.name, + displayName: metadata.displayName, + loaded: Boolean(loaded), + contextSize: metadata.context, + loadedAt: loaded?.loadedAt.toISOString(), + threads: loaded?.threads, + batchSize: loaded?.batchSize + }; + } + + async unloadModel(name: string): Promise { + const loaded = this.models.get(name); + if (!loaded) { + return; + } + safeDispose(loaded.chat); + safeDispose(loaded.context); + safeDispose(loaded.model); + this.models.delete(name); + this.emit("unloaded", name); + } + + ensureLoaded(name: string): ManagedModel { + const loaded = this.models.get(name); + if (!loaded) { + throw new Error(`Model '${name}' is not loaded`); + } + return loaded; + } + + async complete(name: string, prompt: string, overrides?: LoadOverrides): Promise { + const managed = this.ensureLoaded(name); + const options = { ...managed.options, ...overrides }; + const promptTokens = countTokens(managed.context, prompt); + const response: any = await managed.chat.prompt(prompt, { + maxTokens: options.maxTokens, + temperature: options.temperature, + topK: options.topK, + topP: options.topP, + repeatPenalty: options.repeatPenalty + }); + const text = typeof response === "string" ? response : String(response?.output ?? response); + const completionTokens = countTokens(managed.context, text); + return { text, promptTokens, completionTokens }; + } + + async chat( + name: string, + messages: { role: string; content: string }[], + overrides?: LoadOverrides + ): Promise { + const prompt = messages + .map((message) => `${message.role.toUpperCase()}: ${message.content}`) + .join("\n\n"); + return this.complete(name, prompt, overrides); + } + + async embed(name: string, input: string | string[]): Promise { + const managed = this.ensureLoaded(name); + if (typeof managed.context.createEmbedding !== "function") { + throw new Error("Model does not support embeddings on this backend"); + } + const items = Array.isArray(input) ? input : [input]; + const tokens = items.reduce((acc, item) => acc + countTokens(managed.context, item), 0); + const embedding = await managed.context.createEmbedding(items); + return { embedding: Array.isArray(embedding) ? embedding[0] : embedding, tokens }; + } + + resolveModelPath(metadata: ModelMetadata): string { + return path.join(this.config.modelsDir, metadata.name, metadata.file); + } +} + +function safeDispose(target: any): void { + if (target && typeof target.dispose === "function") { + try { + target.dispose(); + } catch (error) { + console.warn("failed to dispose llama resource", error); + } + } +} + +function countTokens(context: any, text: string): number { + if (context && typeof context.tokenize === "function") { + const tokens = context.tokenize(text); + if (Array.isArray(tokens)) { + return tokens.length; + } + } + return Math.max(1, Math.ceil(text.length / 4)); +} + +function requireMetadata(name: string): ModelMetadata { + const metadata = getModelMetadata(name); + if (!metadata) { + throw new Error(`Model '${name}' is not supported`); + } + return metadata; +} diff --git a/apps/llmserver/src/server.ts b/apps/llmserver/src/server.ts new file mode 100644 index 0000000..0b0a0c5 --- /dev/null +++ b/apps/llmserver/src/server.ts @@ -0,0 +1,380 @@ +import "dotenv/config"; +import http from "node:http"; +import express, { Request, Response } from "express"; +import cors from "cors"; +import { Pool } from "pg"; +import { v4 as uuidv4 } from "uuid"; +import expressWs from "express-ws"; +import { z } from "zod"; + +import { loadConfig } from "./config"; +import { ModelManager } from "./modelManager"; +import { TokenTracker } from "./tokenTracker"; +import { requestCounter, tokenCounter, inferenceHistogram, activeSessions, metricsHandler } from "./metrics"; +import { downloadModel, listAvailableDownloads } from "./downloader"; +import { verifyAdminToken } from "./auth"; + +interface UserContext { + userId?: number; + requestId?: string; +} + +const chatSchema = z.object({ + model: z.string(), + messages: z + .array( + z.object({ + role: z.string().min(1), + content: z.string().min(1) + }) + ) + .min(1), + temperature: z.number().min(0).max(2).optional(), + top_k: z.number().int().min(1).max(200).optional(), + top_p: z.number().min(0).max(1).optional(), + repeat_penalty: z.number().min(0).max(2).optional(), + max_tokens: z.number().int().min(1).max(4096).optional() +}); + +const completionSchema = z.object({ + model: z.string(), + prompt: z.string().min(1), + temperature: z.number().min(0).max(2).optional(), + top_k: z.number().int().min(1).max(200).optional(), + top_p: z.number().min(0).max(1).optional(), + repeat_penalty: z.number().min(0).max(2).optional(), + max_tokens: z.number().int().min(1).max(4096).optional() +}); + +const embeddingsSchema = z.object({ + model: z.string(), + input: z.union([z.string(), z.array(z.string().min(1)).min(1)]) +}); + +const adminLoadSchema = z.object({ + model: z.string(), + temperature: z.number().min(0).max(2).optional(), + top_k: z.number().int().min(1).max(200).optional(), + top_p: z.number().min(0).max(1).optional(), + repeat_penalty: z.number().min(0).max(2).optional(), + max_tokens: z.number().int().min(1).max(4096).optional() +}); + +const config = loadConfig(); +const pool = new Pool({ connectionString: config.databaseUrl }); +const tokenTracker = new TokenTracker(pool); +const modelManager = new ModelManager(config); + +const app = express(); +const server = http.createServer(app); +expressWs(app, server); + +app.use(cors()); +app.use(express.json({ limit: "2mb" })); + +app.get("/health", (_req, res) => { + res.json({ status: "ok", uptime: process.uptime() }); +}); + +app.get("/metrics", async (_req, res) => { + const metrics = await metricsHandler(); + res.setHeader("Content-Type", "text/plain; version=0.0.4"); + res.send(metrics); +}); + +app.get("/admin/models", (_req, res) => { + res.json({ models: modelManager.listModels(), downloads: listAvailableDownloads() }); +}); + +app.get("/admin/status", (_req, res) => { + res.json({ + uptime: process.uptime(), + memory: process.memoryUsage(), + models: modelManager.listModels() + }); +}); + +app.post("/admin/download", async (req, res) => { + try { + requireAdmin(req.headers.authorization); + const { model } = z.object({ model: z.string() }).parse(req.body); + const progress = await downloadModel(config, model); + res.json({ status: "downloaded", progress }); + } catch (error) { + respondError(res, error); + } +}); + +app.post("/admin/load", async (req, res) => { + try { + requireAdmin(req.headers.authorization); + const payload = adminLoadSchema.parse(req.body); + const status = await modelManager.loadModel(payload.model, { + temperature: payload.temperature, + topK: payload.top_k, + topP: payload.top_p, + repeatPenalty: payload.repeat_penalty, + maxTokens: payload.max_tokens + }); + res.json({ status: "loaded", model: status }); + } catch (error) { + respondError(res, error); + } +}); + +app.post("/admin/unload", async (req, res) => { + try { + requireAdmin(req.headers.authorization); + const { model } = z.object({ model: z.string() }).parse(req.body); + await modelManager.unloadModel(model); + res.json({ status: "unloaded", model }); + } catch (error) { + respondError(res, error); + } +}); + +app.post("/v1/chat/completions", async (req, res) => { + const context = extractUserContext(req); + const parsed = chatSchema.safeParse(req.body); + if (!parsed.success) { + return res.status(400).json({ error: parsed.error.flatten() }); + } + const payload = parsed.data; + requestCounter.inc({ endpoint: "chat", model: payload.model }); + const stopTimer = inferenceHistogram.startTimer({ endpoint: "chat", model: payload.model }); + activeSessions.inc(); + try { + const result = await modelManager.chat(payload.model, payload.messages, { + temperature: payload.temperature, + topK: payload.top_k, + topP: payload.top_p, + repeatPenalty: payload.repeat_penalty, + maxTokens: payload.max_tokens + }); + await recordTokens(context, payload.model, "chat", result.promptTokens, result.completionTokens); + tokenCounter.inc({ model: payload.model, type: "prompt" }, result.promptTokens); + tokenCounter.inc({ model: payload.model, type: "completion" }, result.completionTokens); + res.json(buildChatResponse(payload.model, result.text, result.promptTokens, result.completionTokens)); + } catch (error) { + respondError(res, error); + } finally { + stopTimer(); + activeSessions.dec(); + } +}); + +app.post("/v1/completions", async (req, res) => { + const context = extractUserContext(req); + const parsed = completionSchema.safeParse(req.body); + if (!parsed.success) { + return res.status(400).json({ error: parsed.error.flatten() }); + } + const payload = parsed.data; + requestCounter.inc({ endpoint: "completion", model: payload.model }); + const stopTimer = inferenceHistogram.startTimer({ endpoint: "completion", model: payload.model }); + try { + const result = await modelManager.complete(payload.model, payload.prompt, { + temperature: payload.temperature, + topK: payload.top_k, + topP: payload.top_p, + repeatPenalty: payload.repeat_penalty, + maxTokens: payload.max_tokens + }); + await recordTokens(context, payload.model, "completion", result.promptTokens, result.completionTokens); + tokenCounter.inc({ model: payload.model, type: "prompt" }, result.promptTokens); + tokenCounter.inc({ model: payload.model, type: "completion" }, result.completionTokens); + res.json({ + id: uuidv4(), + object: "text_completion", + created: Math.floor(Date.now() / 1000), + model: payload.model, + choices: [ + { + index: 0, + text: result.text, + finish_reason: "stop" + } + ], + usage: { + prompt_tokens: result.promptTokens, + completion_tokens: result.completionTokens, + total_tokens: result.promptTokens + result.completionTokens + } + }); + } catch (error) { + respondError(res, error); + } finally { + stopTimer(); + } +}); + +app.post("/v1/embeddings", async (req, res) => { + const context = extractUserContext(req); + const parsed = embeddingsSchema.safeParse(req.body); + if (!parsed.success) { + return res.status(400).json({ error: parsed.error.flatten() }); + } + const payload = parsed.data; + requestCounter.inc({ endpoint: "embeddings", model: payload.model }); + const stopTimer = inferenceHistogram.startTimer({ endpoint: "embeddings", model: payload.model }); + try { + const result = await modelManager.embed(payload.model, payload.input); + await recordTokens(context, payload.model, "embeddings", result.tokens, 0); + tokenCounter.inc({ model: payload.model, type: "prompt" }, result.tokens); + res.json({ + object: "list", + data: [ + { + object: "embedding", + embedding: result.embedding, + index: 0 + } + ], + model: payload.model, + usage: { + prompt_tokens: result.tokens, + completion_tokens: 0, + total_tokens: result.tokens + } + }); + } catch (error) { + respondError(res, error); + } finally { + stopTimer(); + } +}); + +app.ws("/v1/stream", async (ws, req) => { + const context = extractUserContext(req as Request); + ws.on("message", async (raw) => { + try { + const payload = chatSchema.parse(JSON.parse(raw.toString())); + requestCounter.inc({ endpoint: "stream", model: payload.model }); + activeSessions.inc(); + const stopTimer = inferenceHistogram.startTimer({ endpoint: "stream", model: payload.model }); + try { + const result = await modelManager.chat(payload.model, payload.messages, { + temperature: payload.temperature, + topK: payload.top_k, + topP: payload.top_p, + repeatPenalty: payload.repeat_penalty, + maxTokens: payload.max_tokens + }); + await recordTokens(context, payload.model, "stream", result.promptTokens, result.completionTokens); + tokenCounter.inc({ model: payload.model, type: "prompt" }, result.promptTokens); + tokenCounter.inc({ model: payload.model, type: "completion" }, result.completionTokens); + ws.send( + JSON.stringify({ + type: "chunk", + data: result.text, + usage: { + prompt_tokens: result.promptTokens, + completion_tokens: result.completionTokens, + total_tokens: result.promptTokens + result.completionTokens + } + }) + ); + ws.send(JSON.stringify({ type: "done" })); + } catch (error) { + ws.send(JSON.stringify({ type: "error", message: messageFromError(error) })); + } finally { + stopTimer(); + activeSessions.dec(); + } + } catch (error) { + ws.send(JSON.stringify({ type: "error", message: messageFromError(error) })); + } + }); +}); + +server.listen(config.port, config.host, () => { + // eslint-disable-next-line no-console + console.log(`LLM server listening on ${config.host}:${config.port}`); +}); + +async function recordTokens( + context: UserContext, + model: string, + endpoint: string, + promptTokens: number, + completionTokens: number +): Promise { + if (!context.userId) { + return; + } + try { + await tokenTracker.recordUsage({ + userId: context.userId, + model, + endpoint, + promptTokens, + completionTokens, + requestId: context.requestId + }); + } catch (error) { + if (messageFromError(error).includes("Insufficient token")) { + throw Object.assign(new Error("insufficient tokens"), { status: 429 }); + } + throw error; + } +} + +function extractUserContext(req: Request): UserContext { + const userHeader = req.headers["x-user-id"]; + const requestId = req.headers["x-request-id"]; + let userId: number | undefined; + if (typeof userHeader === "string") { + userId = parseInt(userHeader, 10); + } + return { + userId: Number.isFinite(userId) ? userId : undefined, + requestId: typeof requestId === "string" ? requestId : undefined + }; +} + +function requireAdmin(authorization?: string): void { + if (!authorization) { + throw Object.assign(new Error("missing authorization header"), { status: 401 }); + } + const token = authorization.replace(/^Bearer\s+/i, ""); + verifyAdminToken(config, token); +} + +function buildChatResponse(model: string, text: string, promptTokens: number, completionTokens: number) { + return { + id: uuidv4(), + object: "chat.completion", + created: Math.floor(Date.now() / 1000), + model, + choices: [ + { + index: 0, + finish_reason: "stop", + message: { + role: "assistant", + content: text + } + } + ], + usage: { + prompt_tokens: promptTokens, + completion_tokens: completionTokens, + total_tokens: promptTokens + completionTokens + } + }; +} + +function respondError(res: Response, error: unknown): void { + const status = (error as { status?: number }).status ?? 500; + res.status(status).json({ error: messageFromError(error) }); +} + +function messageFromError(error: unknown): string { + if (error instanceof Error) { + return error.message; + } + if (typeof error === "string") { + return error; + } + return "unknown error"; +} diff --git a/apps/llmserver/src/tokenTracker.ts b/apps/llmserver/src/tokenTracker.ts new file mode 100644 index 0000000..d6da247 --- /dev/null +++ b/apps/llmserver/src/tokenTracker.ts @@ -0,0 +1,72 @@ +import { Pool, PoolClient } from "pg"; +import { getModelMetadata, ModelMetadata } from "./catalog"; + +export interface TokenUsage { + userId: number; + model: string; + endpoint: string; + promptTokens: number; + completionTokens: number; + requestId?: string; +} + +export class TokenTracker { + private readonly pool: Pool; + + constructor(pool: Pool) { + this.pool = pool; + } + + async recordUsage(usage: TokenUsage): Promise { + const totalTokens = usage.promptTokens + usage.completionTokens; + if (totalTokens <= 0) { + return; + } + const metadata = getModelMetadata(usage.model); + if (!metadata) { + throw new Error(`Unsupported model '${usage.model}'`); + } + const client = await this.pool.connect(); + try { + await client.query("BEGIN"); + const modelId = await this.ensureModel(client, metadata); + const balanceRow = await client.query( + "SELECT token_balance FROM users WHERE id = $1 FOR UPDATE", + [usage.userId] + ); + if (balanceRow.rowCount === 0) { + throw new Error("User not found for token accounting"); + } + const balance = Number(balanceRow.rows[0].token_balance); + if (balance < totalTokens) { + throw new Error("Insufficient token balance"); + } + await client.query( + "INSERT INTO tokens_used (user_id, model_id, tokens, endpoint) VALUES ($1, $2, $3, $4)", + [usage.userId, modelId, totalTokens, usage.endpoint] + ); + await client.query( + "UPDATE users SET token_balance = token_balance - $1 WHERE id = $2", + [totalTokens, usage.userId] + ); + await client.query("COMMIT"); + } catch (err) { + await client.query("ROLLBACK"); + throw err; + } finally { + client.release(); + } + } + + private async ensureModel(client: PoolClient, metadata: ModelMetadata): Promise { + const existing = await client.query("SELECT id FROM models WHERE name = $1", [metadata.name]); + if (existing.rowCount > 0) { + return Number(existing.rows[0].id); + } + const inserted = await client.query( + "INSERT INTO models (name, huggingface_url, context_size, cost_per_token, is_loaded) VALUES ($1, $2, $3, $4, false) RETURNING id", + [metadata.name, `https://huggingface.co/${metadata.huggingFaceRepo}`, metadata.context, metadata.costPerToken] + ); + return Number(inserted.rows[0].id); + } +} diff --git a/apps/llmserver/tsconfig.json b/apps/llmserver/tsconfig.json new file mode 100644 index 0000000..9d77330 --- /dev/null +++ b/apps/llmserver/tsconfig.json @@ -0,0 +1,17 @@ +{ + "compilerOptions": { + "target": "ES2020", + "module": "commonjs", + "moduleResolution": "node", + "rootDir": "src", + "outDir": "dist", + "strict": true, + "esModuleInterop": true, + "forceConsistentCasingInFileNames": true, + "skipLibCheck": true, + "resolveJsonModule": true, + "sourceMap": true + }, + "include": ["src/**/*.ts"], + "exclude": ["node_modules", "dist"] +} diff --git a/apps/studio-ui/.eslintrc.cjs b/apps/studio-ui/.eslintrc.cjs new file mode 100644 index 0000000..84ffbfe --- /dev/null +++ b/apps/studio-ui/.eslintrc.cjs @@ -0,0 +1,23 @@ +module.exports = { + root: true, + env: { + browser: true, + es2021: true + }, + extends: ['eslint:recommended', 'plugin:react/recommended', 'prettier'], + parserOptions: { + ecmaVersion: 'latest', + sourceType: 'module' + }, + settings: { + react: { + version: 'detect' + } + }, + plugins: ['react'], + rules: { + 'react/prop-types': 'off', + 'react/jsx-uses-react': 'off', + 'react/react-in-jsx-scope': 'off' + } +}; diff --git a/apps/studio-ui/index.html b/apps/studio-ui/index.html new file mode 100644 index 0000000..1f8eb94 --- /dev/null +++ b/apps/studio-ui/index.html @@ -0,0 +1,12 @@ + + + + + + CyberDevStudio + + +
+ + + diff --git a/apps/studio-ui/package.json b/apps/studio-ui/package.json new file mode 100644 index 0000000..083407f --- /dev/null +++ b/apps/studio-ui/package.json @@ -0,0 +1,41 @@ +{ + "name": "studio-ui", + "version": "1.0.0", + "private": true, + "type": "module", + "scripts": { + "dev": "vite", + "build": "vite build", + "preview": "vite preview", + "lint": "eslint --ext .ts,.tsx src", + "test": "vitest" + }, + "dependencies": { + "@monaco-editor/react": "^4.6.0", + "@tanstack/react-query": "^5.29.0", + "@xterm/xterm": "^5.3.0", + "clsx": "^2.1.0", + "react": "^18.2.0", + "react-dom": "^18.2.0", + "react-router-dom": "^6.22.3", + "socket.io-client": "^4.7.5" + }, + "devDependencies": { + "@testing-library/jest-dom": "^6.4.2", + "@testing-library/react": "^14.2.1", + "@testing-library/user-event": "^14.6.1", + "@types/react": "^18.2.48", + "@types/react-dom": "^18.2.18", + "@vitejs/plugin-react": "^4.3.1", + "autoprefixer": "^10.4.18", + "eslint": "^8.56.0", + "eslint-config-prettier": "^9.1.0", + "eslint-plugin-react": "^7.33.2", + "jsdom": "^24.0.0", + "postcss": "^8.4.35", + "tailwindcss": "^3.4.1", + "typescript": "^5.3.3", + "vite": "^5.1.4", + "vitest": "^1.4.0" + } +} diff --git a/apps/studio-ui/postcss.config.cjs b/apps/studio-ui/postcss.config.cjs new file mode 100644 index 0000000..5cbc2c7 --- /dev/null +++ b/apps/studio-ui/postcss.config.cjs @@ -0,0 +1,6 @@ +module.exports = { + plugins: { + tailwindcss: {}, + autoprefixer: {} + } +}; diff --git a/apps/studio-ui/src/App.test.tsx b/apps/studio-ui/src/App.test.tsx new file mode 100644 index 0000000..089464b --- /dev/null +++ b/apps/studio-ui/src/App.test.tsx @@ -0,0 +1,24 @@ +import { describe, expect, it } from 'vitest'; +import { render, screen } from '@testing-library/react'; +import { MemoryRouter } from 'react-router-dom'; +import App from './App'; + +function renderApp(initialEntries: string[] = ['/']) { + render( + + + + ); +} + +describe('App', () => { + it('renders the editor tab by default', () => { + renderApp(); + expect(screen.getByText('Save')).toBeInTheDocument(); + }); + + it('allows navigation to LLM tab', async () => { + renderApp(['/llm']); + expect(await screen.findByText(/Generate/)).toBeInTheDocument(); + }); +}); diff --git a/apps/studio-ui/src/App.tsx b/apps/studio-ui/src/App.tsx new file mode 100644 index 0000000..7a2bb1c --- /dev/null +++ b/apps/studio-ui/src/App.tsx @@ -0,0 +1,62 @@ +import { useEffect, useState } from 'react'; +import { Outlet, Route, Routes, useLocation } from 'react-router-dom'; +import { Sidebar } from './components/Sidebar'; +import { EditorView } from './components/EditorView'; +import { AgentChat } from './components/AgentChat'; +import { ExecutionView } from './components/ExecutionView'; +import { TerminalView } from './components/TerminalView'; +import { LLMPlayground } from './components/LLMPlayground'; +import { AdminPanel } from './components/AdminPanel'; +import { TopBar } from './components/TopBar'; +import { StudioContextProvider } from './hooks/useStudioContext'; + +const tabs = [ + { name: 'Code', path: '/code' }, + { name: 'Chat', path: '/chat' }, + { name: 'Run', path: '/run' }, + { name: 'Terminal', path: '/terminal' }, + { name: 'LLM', path: '/llm' }, + { name: 'Admin', path: '/admin', protected: true } +] as const; + +type Tab = (typeof tabs)[number]; + +function Layout() { + const location = useLocation(); + const [activeTab, setActiveTab] = useState('/code'); + + useEffect(() => { + const normalized = location.pathname === '/' ? '/code' : (location.pathname as Tab['path']); + setActiveTab(normalized); + }, [location.pathname]); + + return ( +
+ +
+ +
+ +
+
+
+ ); +} + +export default function App() { + return ( + + + }> + } /> + } /> + } /> + } /> + } /> + } /> + } /> + + + + ); +} diff --git a/apps/studio-ui/src/api/rpc.ts b/apps/studio-ui/src/api/rpc.ts new file mode 100644 index 0000000..0d3d2e0 --- /dev/null +++ b/apps/studio-ui/src/api/rpc.ts @@ -0,0 +1,94 @@ +export interface JsonRpcRequest

{ + jsonrpc: '2.0'; + id: string; + method: string; + params?: P; +} + +export interface JsonRpcSuccess { + jsonrpc: '2.0'; + id: string; + result: R; +} + +export interface JsonRpcError { + code: number; + message: string; + data?: unknown; +} + +export interface JsonRpcFailure { + jsonrpc: '2.0'; + id: string; + error: JsonRpcError; +} + +export type JsonRpcResponse = JsonRpcSuccess | JsonRpcFailure; + +export class RpcError extends Error { + public readonly code: number; + public readonly data?: unknown; + + constructor(message: string, code: number, data?: unknown) { + super(message); + this.name = 'RpcError'; + this.code = code; + this.data = data; + } +} + +export class RpcClient { + private readonly endpoint: string; + private token?: string; + private apiKey?: string; + + constructor(endpoint: string, token?: string, apiKey?: string) { + this.endpoint = endpoint; + this.token = token; + this.apiKey = apiKey; + } + + setBearerToken(token: string | undefined) { + this.token = token; + } + + setApiKey(apiKey: string | undefined) { + this.apiKey = apiKey; + } + + async call(method: string, params?: P): Promise { + const request: JsonRpcRequest

= { + jsonrpc: '2.0', + id: crypto.randomUUID(), + method, + params + }; + + const headers: HeadersInit = { + 'Content-Type': 'application/json' + }; + if (this.token) { + headers['Authorization'] = `Bearer ${this.token}`; + } + if (this.apiKey) { + headers['X-API-Key'] = this.apiKey; + } + + const response = await fetch(this.endpoint, { + method: 'POST', + headers, + body: JSON.stringify(request) + }); + + if (!response.ok) { + throw new RpcError(`RPC transport failure: ${response.statusText}`, response.status); + } + + const payload = (await response.json()) as JsonRpcResponse; + if ('error' in payload) { + throw new RpcError(payload.error.message, payload.error.code, payload.error.data); + } + + return payload.result; + } +} diff --git a/apps/studio-ui/src/components/AdminPanel.tsx b/apps/studio-ui/src/components/AdminPanel.tsx new file mode 100644 index 0000000..f4fe55e --- /dev/null +++ b/apps/studio-ui/src/components/AdminPanel.tsx @@ -0,0 +1,193 @@ +import { useEffect, useMemo, useState } from 'react'; +import { useStudioContext } from '../hooks/useStudioContext'; + +interface ModelInfo { + id: string; + name: string; + status: 'available' | 'loaded'; + huggingface_url?: string; + context_size: number; + memory_mb?: number; +} + +interface UserInfo { + id: number; + username: string; + role: string; + token_balance: number; +} + +interface TokenUsage { + timestamp: string; + tokens: number; + model: string; +} + +export function AdminPanel() { + const { rpc, profile } = useStudioContext(); + const [models, setModels] = useState([]); + const [users, setUsers] = useState([]); + const [usage, setUsage] = useState([]); + const [loading, setLoading] = useState(false); + const [error, setError] = useState(null); + + const canAccess = profile.role === 'admin'; + + const loadData = async () => { + setLoading(true); + setError(null); + try { + const [modelsResponse, usersResponse, usageResponse] = await Promise.all([ + rpc.call<{ models: ModelInfo[] }>('llm.list_models'), + rpc.call<{ users: UserInfo[] }>('admin.users.list'), + rpc.call<{ usage: TokenUsage[] }>('admin.tokens.recent') + ]); + setModels(modelsResponse.models); + setUsers(usersResponse.users); + setUsage(usageResponse.usage); + } catch (err) { + if (err instanceof Error) { + setError(err.message); + } else { + setError('Unable to load admin data'); + } + } finally { + setLoading(false); + } + }; + + useEffect(() => { + if (canAccess) { + loadData().catch((err) => console.error(err)); + } + }, [canAccess]); + + const handleModelAction = async (model: ModelInfo, action: 'load' | 'unload' | 'download') => { + setLoading(true); + setError(null); + try { + if (action === 'download') { + await rpc.call('llm.download', { id: model.id }); + } else if (action === 'load') { + await rpc.call('llm.start', { id: model.id, options: { temperature: 0.2, top_k: 40, top_p: 0.95 } }); + } else if (action === 'unload') { + await rpc.call('llm.stop', { id: model.id }); + } + await loadData(); + } catch (err) { + if (err instanceof Error) { + setError(err.message); + } else { + setError('Unable to update model state'); + } + } finally { + setLoading(false); + } + }; + + const totalTokens = useMemo(() => usage.reduce((sum, entry) => sum + entry.tokens, 0), [usage]); + + if (!canAccess) { + return ( +

+ Admin access required. +
+ ); + } + + return ( +
+
+
+

Admin Control Center

+ +
+ {loading &&

Refreshing data…

} + {error &&

{error}

} +
+
+
+
+

Models

+ {models.length} available +
+
+ {models.map((model) => ( +
+
+
+

{model.name}

+

Context {model.context_size} · Status {model.status}

+
+
+ + + +
+
+ {model.huggingface_url && ( +

Source: {model.huggingface_url}

+ )} +
+ ))} +
+
+
+
+

Users

+ {users.length} accounts +
+
+ {users.map((user) => ( +
+
+
+

{user.username}

+

Role {user.role}

+
+ Tokens {user.token_balance} +
+
+ ))} +
+
+
+
+

Recent Token Usage

+ Total tokens {totalTokens} +
+
+ {usage.map((entry) => ( +
+ {entry.model} + {new Date(entry.timestamp).toLocaleString()} + {entry.tokens} tokens +
+ ))} + {usage.length === 0 &&

No usage records available.

} +
+
+
+
+ ); +} diff --git a/apps/studio-ui/src/components/AgentChat.tsx b/apps/studio-ui/src/components/AgentChat.tsx new file mode 100644 index 0000000..f633a2b --- /dev/null +++ b/apps/studio-ui/src/components/AgentChat.tsx @@ -0,0 +1,171 @@ +import { FormEvent, useEffect, useRef, useState } from 'react'; +import { io, Socket } from 'socket.io-client'; +import { useStudioContext } from '../hooks/useStudioContext'; + +interface ChatMessage { + id: string; + role: 'user' | 'agent' | 'system'; + agent?: string; + content: string; + createdAt: string; +} + +const agentOptions = [ + { id: 'code', label: 'CodeAgent' }, + { id: 'test', label: 'TestAgent' }, + { id: 'design', label: 'DesignAgent' }, + { id: 'debug', label: 'DebugAgent' }, + { id: 'security', label: 'SecurityAgent' }, + { id: 'doc', label: 'DocAgent' } +]; + +export function AgentChat() { + const { rpc } = useStudioContext(); + const [messages, setMessages] = useState([]); + const [agent, setAgent] = useState(agentOptions[0].id); + const [input, setInput] = useState(''); + const [isStreaming, setIsStreaming] = useState(false); + const socketRef = useRef(null); + const viewportRef = useRef(null); + + useEffect(() => { + const socket = io(import.meta.env.VITE_WS_URL ?? '/ws', { + transports: ['websocket'], + reconnection: true, + reconnectionAttempts: 5 + }); + socketRef.current = socket; + socket.on('agent-message', (payload: ChatMessage) => { + setMessages((current) => [...current, payload]); + }); + socket.on('connect', () => { + console.debug('Agent chat connected'); + }); + socket.on('disconnect', () => { + console.debug('Agent chat disconnected'); + }); + return () => { + socket.disconnect(); + }; + }, []); + + useEffect(() => { + if (viewportRef.current) { + viewportRef.current.scrollTop = viewportRef.current.scrollHeight; + } + }, [messages]); + + const handleSubmit = async (event: FormEvent) => { + event.preventDefault(); + if (!input.trim()) { + return; + } + + const userMessage: ChatMessage = { + id: crypto.randomUUID(), + role: 'user', + content: input, + agent, + createdAt: new Date().toISOString() + }; + setMessages((current) => [...current, userMessage]); + setInput(''); + setIsStreaming(true); + + try { + const response = await rpc.call<{ task_id: string }>('agent.dispatch', { + agent, + prompt: userMessage.content, + metadata: { + source: 'studio-ui', + request_id: userMessage.id + } + }); + setMessages((current) => [ + ...current, + { + id: response.task_id, + role: 'system', + content: `Dispatched task ${response.task_id}`, + createdAt: new Date().toISOString(), + agent + } + ]); + } catch (error) { + setMessages((current) => [ + ...current, + { + id: crypto.randomUUID(), + role: 'system', + content: error instanceof Error ? error.message : 'Failed to dispatch agent task', + createdAt: new Date().toISOString(), + agent + } + ]); + } finally { + setIsStreaming(false); + } + }; + + return ( +
+
+
+ + +
+ + {isStreaming ? 'Awaiting agent response…' : 'Idle'} + +
+
+ {messages.map((message) => ( +
+
+ + {message.role.toUpperCase()} {message.agent ? `· ${message.agent}` : ''} + + {new Date(message.createdAt).toLocaleTimeString()} +
+

{message.content}

+
+ ))} +
+
+
+