summaryrefslogtreecommitdiff
path: root/server/src/ai.rs
diff options
context:
space:
mode:
authorSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
committerSomhairle H. Marisol <[email protected]>2026-09-17 14:32:37 +0800
commit5c0ba37eda80d39e6ceca59bb1d5f4942f858995 (patch)
tree948723f9cedf7ccb0707fa6ee516bd30fe20fd10 /server/src/ai.rs
downloadstrategy-lab-5c0ba37eda80d39e6ceca59bb1d5f4942f858995.tar.gz
chore: establish Strategy Lab source baseline (development, not release)
Diffstat (limited to 'server/src/ai.rs')
-rw-r--r--server/src/ai.rs319
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");
+ }
+}