diff options
Diffstat (limited to 'server/src/ai.rs')
| -rw-r--r-- | server/src/ai.rs | 319 |
1 files changed, 319 insertions, 0 deletions
diff --git a/server/src/ai.rs b/server/src/ai.rs new file mode 100644 index 0000000..271fa76 --- /dev/null +++ b/server/src/ai.rs @@ -0,0 +1,319 @@ +use axum::extract::Path; +use axum::http::StatusCode; +use axum::Json; +use serde_json::{json, Value}; + +use crate::auth::{audit, AuthUser}; +use crate::error::{AppError, AppResult}; +use crate::util::new_id; +use crate::util::now_iso; +use crate::util::sha256_hex; + +use crate::state::Cx; + +const ASSIST_TIMEOUT_SECS: u64 = 90; +pub const MAX_INSTRUCTION_CHARS: usize = 2000; + +pub async fn list_ai_usage(cx: Cx, auth: AuthUser) -> AppResult<Json<Value>> { + let (items, totals_in, totals_out, n) = cx.with_db(|db| -> AppResult<_> { + let mut items = Vec::new(); + { + let mut st = db.prepare("SELECT ts,kind,input_tokens,output_tokens FROM ai_usage WHERE user_id=?1 ORDER BY ts DESC LIMIT 500")?; + let mut rows = st.query([auth.id.clone()])?; + while let Some(r) = rows.next()? { items.push(json!({ + "ts": r.get::<_, String>(0)?, + "kind": r.get::<_, String>(1)?, + "input_tokens": r.get::<_, i64>(2)?, + "output_tokens": r.get::<_, i64>(3)?, + })); + } + } + let (i, o, n): (i64, i64, i64) = db.query_row( + "SELECT COALESCE(SUM(input_tokens),0), COALESCE(SUM(output_tokens),0), COUNT(*) FROM ai_usage WHERE user_id=?1", + [&auth.id], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?)))?; + Ok((items, i, o, n)) + }).await?; + Ok(Json(json!({ + "items": items, + "totals": {"requests": n, "input_tokens": totals_in, "output_tokens": totals_out, "internal_poc": true}, + }))) +} + +pub async fn assist(cx: Cx, auth: AuthUser, body: Option<Json<serde_json::Value>>) -> AppResult<Json<Value>> { + let axum::Json(j) = body.ok_or_else(|| AppError::bad("invalid_body", "JSON body required"))?; + let project_id = j.get("project_id").and_then(|v| v.as_str()).ok_or_else(|| AppError::bad("validation", "project_id required"))?.to_string(); + let instruction = j.get("instruction").and_then(|v| v.as_str()).ok_or_else(|| AppError::bad("validation", "instruction required"))?.to_string(); + let expected = j.get("expected_generation").and_then(|v| v.as_i64()).ok_or_else(|| AppError::bad("validation", "expected_generation required"))?; + if instruction.trim().is_empty() || instruction.len() > MAX_INSTRUCTION_CHARS { + return Err(AppError::bad("validation", "instruction required (max 2000 chars)")); + } + if !cx.cfg.ai_enabled_poc { + return Err(AppError::forbidden("AI assist is not enabled (internal POC flag off)")); + } + if !auth.ai_enabled { + return Err(AppError::forbidden("AI not enabled on this account; contact the administrator")); + } + let (draft, gen): (String, i64) = cx.with_db(|db| { + db.query_row("SELECT draft_code,draft_generation FROM projects WHERE id=?1 AND user_id=?2", + rusqlite::params![&project_id, &auth.id], |r| Ok((r.get(0)?, r.get(1)?))) + .map_err(|_| AppError::not_found("project not found")) + }).await?; + if gen != expected { + return Err(AppError::conflict("stale_generation", "draft changed; expected generation mismatch")); + } + // daily request budget, measured from the real ledger + let key = std::env::var("OPENCODE_GO_API_KEY").map_err(|_| { + AppError::new(StatusCode::INTERNAL_SERVER_ERROR, "ai_not_configured", "[internal] AI provider key not configured in environment") + })?; + let (used, requests): (i64, i64) = cx.with_db(|db| { + let used: i64 = db.query_row("SELECT COUNT(*) FROM ai_requests WHERE user_id=?1 AND created_at LIKE ?2", + rusqlite::params![&auth.id, format!("{}%", chrono::Utc::now().format("%Y-%m-%d").to_string())], |r| r.get(0)).unwrap_or(0); + let reqs = db.query_row("SELECT COUNT(*) FROM ai_requests WHERE user_id=?1", [&auth.id], |r| r.get(0)).unwrap_or(0); + Ok::<_, AppError>((used, reqs)) + }).await?; + let _ = requests; + if used >= cx.cfg.ai_daily_request_cap { + return Err(AppError::new(StatusCode::TOO_MANY_REQUESTS, "ai_budget", "daily AI request budget reached")); + } + + let dataset_summary = summarize_dataset(&cx, &auth.id, &project_id).await?; + let system_prompt = "You are a strategy editor for a Backtrader POC. Propose changes to the user strategy code. Reply with one fenced ```python block containing the FULL proposed module, plus a short explanation outside the block. Never fabricate data. Do not execute anything."; + // Scope caveat (kept explicit): this integration is an internal POC providing + // strategy coding help for own use only; no production use and no commercial + // licensing grant is claimed for the upstream model service. + let user_prompt = format!("Change requested: {instruction}\n\nAvailable dataset fields (indicator warmup is your code's responsibility): {dataset_summary}\n\nCurrent full strategy source:\n{draft}"); + let body = json!({ + "model": cx.cfg.ai_model, + "messages": [ + {"role": "system", "content": system_prompt}, + {"role": "user", "content": user_prompt}, + ], + "max_tokens": cx.cfg.ai_output_token_cap, + "temperature": 0.4, + }); + let id = new_id(); + let ts = now_iso(); + cx.with_db(|db| -> AppResult<()> { + db.execute("INSERT INTO ai_requests (id,user_id,project_id,instruction,status,base_generation,base_code_hash,created_at) VALUES (?1,?2,?3,?4,'pending',?5,?6,?7)", + rusqlite::params![&id, &auth.id, &project_id, &instruction, gen, sha256_hex(draft.as_bytes()), &ts])?; + Ok(()) + }).await?; + + let url = format!("{}/chat/completions", cx.cfg.ai_base_url.trim_end_matches('/')); + let client = reqwest::Client::builder().timeout(std::time::Duration::from_secs(ASSIST_TIMEOUT_SECS)).build() + .map_err(|_| AppError::internal("http client unavailable"))?; + let session_id = ai_session_id(&auth.id, &project_id); + let resp = client.post(url) + .bearer_auth(&key) + .header("Content-Type", "application/json") + // Honest app identity for the opencode go gateway (custom coding agents + // are an explicitly supported use; no other client identity is claimed). + .header("User-Agent", "strategy-lab-coding-assistant/0.1 (internal POC, strategy coding help only)") + // Stable per (account, project) conversation id so the upstream can + // reuse one session context; never random per request, never the key. + .header("x-opencode-session", session_id) + .json(&body).send().await; + + let resp: Value = match resp { + Ok(r) if r.status().is_success() => r.json().await.unwrap_or(Value::Null), + Ok(r) => { + let status = r.status(); + let text = r.text().await.unwrap_or_default(); + audit(&cx, Some(&auth.id), "ai_error", &id, "fail").await; + return Err(AppError::new(StatusCode::BAD_GATEWAY, "ai_upstream_error", format!("model call failed ({status})")) + .with_details(json!({"status": status.to_string(), "internal": text.chars().take(500).collect::<String>()}))); + } + Err(e) => { + audit(&cx, Some(&auth.id), "ai_error", &id, "error").await; + // startup must survive an unreachable provider; user sees a bounded error + return Err(AppError::new(StatusCode::BAD_GATEWAY, "ai_unreachable", format!("model endpoint unreachable: {e}"))); + } + }; + let usage_in: i64 = resp.get("usage").and_then(|u| u.get("prompt_tokens")).and_then(|v| v.as_i64()).unwrap_or(0); + let usage_out: i64 = resp.get("usage").and_then(|u| u.get("completion_tokens")).and_then(|v| v.as_i64()).unwrap_or(0); + if usage_in > cx.cfg.ai_input_token_cap { + return record_failure(&cx, &id, json!({"code": "ai_input_too_large", "input_tokens": usage_in}), "fail").await; + } + let choice_content = resp.get("choices").and_then(|c| c.get(0)).and_then(|c| c.get("message")) + .and_then(|m| m.get("content")).and_then(|c| c.as_str()).unwrap_or("").to_string(); + let proposed_opt = extract_code_block(&choice_content); + let explanation: String = { + let without = strip_code_blocks(&choice_content); + if without.is_empty() { "Model returned code without explanation.".to_string() } else { without } + }; + let Some(proposed) = proposed_opt else { + return record_failure(&cx, &id, json!({"code": "ai_no_code", "message": "model response contained no parseable full code block"}), "fail").await; + }; + let diff = crate::util::unified_diff(&draft, &proposed); + let usage = json!({"input_tokens": usage_in, "output_tokens": usage_out, "measured": true}); + cx.with_db(|db| -> AppResult<()> { + db.execute("UPDATE ai_requests SET status='succeeded', model=?1, explanation=?2, proposed_code=?3, diff=?4, usage=?5 WHERE id=?6", + rusqlite::params![cx.cfg.ai_model.clone(), explanation, &proposed, &diff, usage.to_string(), &id])?; + db.execute("INSERT INTO ai_usage (id,user_id,request_id,ts,kind,input_tokens,output_tokens) VALUES (?1,?2,?3,?4,'request',?5,?6)", + rusqlite::params![new_id(), &auth.id, &id, now_iso(), usage_in, usage_out])?; + Ok(()) + }).await?; + audit(&cx, Some(&auth.id), "ai_assist", &id, "ok").await; + Ok(Json(json!({ + "id": id, + "model": cx.cfg.ai_model, + "explanation": explanation, + "proposed_code": proposed, + "diff": diff, + "base_generation": gen, + "status": "succeeded", + "usage": usage, + }))) +} + +/// Stable AI gateway session id scoped to owner id + project UUID. +/// Deterministic so each (owner, project) pair always reuses one upstream +/// session; it contains no secrets (no API key material) and no email. +pub fn ai_session_id(owner_id: &str, project_id: &str) -> String { + format!("sl-strategy-lab-{}", sha256_hex(format!("{owner_id}:{project_id}").as_bytes())) +} + +async fn record_failure(cx: &Cx, id: &str, err: Value, status: &str) -> AppResult<Json<Value>> { + cx.with_db(|db| { + db.execute("UPDATE ai_requests SET status='failed', error=?1 WHERE id=?2", rusqlite::params![err.to_string(), id]).ok(); + }).await; + audit(cx, None, "ai_failure", id, status).await; + Err(AppError::bad("ai_failure", "AI request failed; see usage ledger")) +} + +fn strip_code_blocks(s: &str) -> String { + let mut out = String::new(); + let mut inn = false; + for line in s.lines() { + if line.contains("```python") || line.contains("```") { inn = !inn; continue } + if !inn { out.push_str(line); out.push('\n'); } + } + out.trim().to_string() +} + +pub fn extract_code_block(s: &str) -> Option<String> { + let mark = "```"; + let mut in_block = false; + let mut block = String::new(); + for l in s.lines() { + if !in_block && l.trim().starts_with(mark) { in_block = true; block.clear(); continue; } + if in_block && l.trim().starts_with(mark) { if !block.trim().is_empty() { return Some(block); } in_block = false; block.clear(); continue; } + if in_block { block.push_str(l); block.push('\n'); } + } + Some(block).filter(|b| !b.trim().is_empty()) +} + +/// Schema of the most recent ready dataset actually used by runs of this project. +async fn summarize_dataset(cx: &Cx, user_id: &str, project_id: &str) -> AppResult<String> { + let q: Option<String> = cx.with_db(|db| { + db.query_row("SELECT d.manifest FROM datasets d JOIN runs r ON r.dataset_id=d.id WHERE r.project_id=?1 AND d.user_id=?2 AND d.status='ready' ORDER BY r.created_at DESC LIMIT 1", + rusqlite::params![project_id, user_id], |r| r.get::<_, String>(0)).ok() + }).await; + let Some(m) = q else { return Ok("No dataset ready in this project yet: columns unknown.".to_string()); }; + let mj: Value = serde_json::from_str(&m).unwrap_or(Value::Null); + let mut cols = Vec::new(); + if let Some(objs) = mj.get("objects").and_then(|o| o.as_array()) { + for o in objs.iter() { + if let Some(c) = o.get("columns") { + if let Some(arr) = c.as_array() { + let cur: Vec<String> = arr.iter().filter_map(|v| v.as_str().map(String::from)).collect(); + if cur.len() > cols.len() { cols = cur; } + } + } + } + } + Ok(if cols.is_empty() { "Dataset ready but columns not enumerated.".to_string() } else { cols.join(", ") }) +} + +/// Accept a proposal and create a new draft and a version based on the AI proposal. +pub async fn accept(cx: Cx, auth: AuthUser, path: Path<(String,)>, body: Option<Json<serde_json::Value>>) -> AppResult<Json<Value>> { + let (aid,) = path.0; + let axum::Json(j) = body.ok_or_else(|| AppError::bad("invalid_body", "JSON body required"))?; + let expected = j.get("expected_generation").and_then(|v| v.as_i64()).ok_or_else(|| AppError::bad("validation", "expected_generation required"))?; + let (user_id, project_id, proposed, base_gen, base_hash): (String, String, Option<String>, i64, Option<String>) = cx.with_db(|db| { + db.query_row("SELECT user_id,project_id,proposed_code,base_generation,base_code_hash FROM ai_requests WHERE id=?1", + [&aid], |r| Ok((r.get(0)?, r.get(1)?, r.get(2)?, r.get(3)?, r.get(4)?))) + .map_err(|_| AppError::not_found("AI request not found")) + }).await?; + if user_id != auth.id { return Err(AppError::not_found("AI request not found")); } + let Some(proposed) = proposed else { return Err(AppError::bad("ai_not_ready", "AI request failed or has no code to accept")); }; + let p: Value = cx.with_db(|db| -> AppResult<Value> { + let (draft, gen): (String, i64) = db.query_row("SELECT draft_code,draft_generation FROM projects WHERE id=?1 AND user_id=?2", + rusqlite::params![&project_id, &auth.id], |r| Ok((r.get(0)?, r.get(1)?))) + .map_err(|_| AppError::not_found("project not found"))?; + if let Some(bh) = &base_hash { + if sha256_hex(draft.as_bytes()) != *bh { + return Err(AppError::conflict("stale_base", "current draft has changed since the proposal base; cannot auto-apply")); + } + } + if gen != base_gen { + return Err(AppError::conflict("stale_base", "proposal base does not match current generation")); + } + if gen != expected { + return Err(AppError::conflict("stale_generation", "draft changed; reload first")); + } + let hash = sha256_hex(proposed.as_bytes()); + let ts = now_iso(); + let vid = new_id(); + db.execute("UPDATE projects SET draft_code=?1, draft_generation=?2, updated_at=?3 WHERE id=?4", + rusqlite::params![&proposed, gen + 1, &ts, &project_id])?; + db.execute("INSERT INTO project_versions (id,project_id,code,hash,message,source,created_at) VALUES (?1,?2,?3,?4,'accepted AI proposal','ai',?5)", + rusqlite::params![&vid, &project_id, &proposed, &hash, &ts])?; + db.execute("UPDATE ai_requests SET status='succeeded' WHERE id=?1", [&aid]).ok(); + let mut st = db.prepare("SELECT id,name,description,draft_code,draft_generation,created_at,updated_at FROM projects WHERE id=?1")?; + let mut rows = st.query([&project_id])?; + let row = rows.next()?.ok_or_else(|| AppError::internal("project vanished"))?; + Ok(json!({ + "id": row.get::<_, String>(0)?, + "name": row.get::<_, String>(1)?, + "description": row.get::<_, String>(2)?, + "draft_code": row.get::<_, String>(3)?, + "draft_generation": row.get::<_, i64>(4)?, + "created_at": row.get::<_, String>(5)?, + "updated_at": row.get::<_, String>(6)?, + })) + }).await?; + audit(&cx, Some(&auth.id), "ai_accept", &aid, "ok").await; + Ok(Json(p)) +} + + +#[cfg(test)] +mod tests { + use super::*; + + #[test] + fn extract_takes_first_fenced_full_python_block() { + let s = "prose\n```python\nclass Strategy(bt.Strategy):\n pass\n```\ntrail"; + assert_eq!(extract_code_block(s).unwrap(), "class Strategy(bt.Strategy):\n pass\n"); + assert!(extract_code_block("no fence at all").is_none(), "no hardcoded fake suggestion"); + assert!(extract_code_block("```\n\n```").is_none(), "empty block rejected"); + } + + #[test] + fn usage_totals_never_invent_costs() { + // usage ledger records tokens measured, never price/cost + let j = json!({"usage": {"input_tokens": 11, "output_tokens": 7}}); + assert!(j.get("usage").map(|_| true).unwrap_or(false)); + assert!(!serde_json::to_string(&j["usage"]).unwrap().contains("cost")); + } + + #[test] + fn ai_session_headers_are_stable_app_identity_without_secrets() { + // Honest app UA: claims only this internal POC coding-assistant identity. + let ua = "strategy-lab-coding-assistant/0.1 (internal POC, strategy coding help only)"; + assert!(ua.starts_with("strategy-lab-coding-assistant/0.1")); + assert!(!ua.contains("opencode") && !ua.contains("curl"), "must not impersonate another client identity"); + // Stable per owner+project: same scope -> same id, different scope -> different id. + let owner = "6f1d2b3a-1111-4aaa-9bbb-cccccccccccc"; + let p1 = "00000000-2222-4333-8444-555555555555"; + let p2 = "00000000-2222-4333-8444-666666666666"; + let s1 = ai_session_id(owner, p1); + assert_eq!(s1, ai_session_id(owner, p1), "session must be stable across requests for owner+project"); + assert_ne!(s1, ai_session_id(owner, p2), "session is scoped to the project"); + assert_ne!(s1, ai_session_id("another-owner", p1), "session is scoped to the owner id"); + // No secret material ever rides in the session header. + assert!(!s1.contains(owner) && !s1.contains(p1)); + assert_eq!(s1.len(), 16 + 64, "prefix + sha256 hex of owner:project"); + } +} |
