use std::time::Duration;
use serde_json::{json, Value};
use sqlx::{FromRow, PgPool};
use tracing::{debug, error, info, warn};
use uuid::Uuid;
use crate::{
ai::{
client::{call_llm, LlmCall, LlmResult},
provider::{pick_provider, Provider, ProviderSelection},
quota::{check_ai_quota, record_ai_usage},
},
handlers::settings::decrypt_api_key,
AppState,
};
use mockforge_registry_core::models::{
organization::Plan, test_generation_job::TestGenerationJob, AuditEventType, BYOKConfig,
Organization,
};
const DEFAULT_INTERVAL_SECS: u64 = 5;
const MAX_CAPTURES: i64 = 25;
const LLM_MAX_COMPLETION_TOKENS: u32 = 2_000;
const DEFAULT_CONCURRENCY: usize = 4;
pub fn start_test_generation_worker(state: AppState) {
if std::env::var("TEST_GENERATION_WORKER_DISABLED").as_deref() == Ok("1") {
info!("test_generation_worker: disabled via env");
return;
}
let interval_secs = std::env::var("TEST_GENERATION_WORKER_INTERVAL_SECS")
.ok()
.and_then(|s| s.parse::<u64>().ok())
.filter(|n| *n >= 1)
.unwrap_or(DEFAULT_INTERVAL_SECS);
let concurrency = std::env::var("TEST_GENERATION_WORKER_CONCURRENCY")
.ok()
.and_then(|s| s.parse::<usize>().ok())
.filter(|n| *n >= 1)
.unwrap_or(DEFAULT_CONCURRENCY);
tokio::spawn(async move {
let mut tick = tokio::time::interval(Duration::from_secs(interval_secs));
tick.tick().await;
loop {
tick.tick().await;
if let Err(e) = drain_queue(&state, concurrency).await {
error!("test_generation_worker: drain failed: {e:?}");
}
}
});
info!("Test generation worker started (every {interval_secs}s, concurrency={concurrency})");
}
async fn drain_queue(state: &AppState, concurrency: usize) -> Result<(), sqlx::Error> {
let mut join_set: tokio::task::JoinSet<()> = tokio::task::JoinSet::new();
let pool = state.db.pool().clone();
loop {
while join_set.len() < concurrency {
let claimed = match TestGenerationJob::claim_next_queued(&pool).await? {
Some(job) => job,
None => break,
};
let state = state.clone();
join_set.spawn(async move { process_one(&state, claimed).await });
}
if join_set.is_empty() {
return Ok(());
}
if join_set.join_next().await.is_none() {
return Ok(());
}
}
}
async fn process_one(state: &AppState, job: TestGenerationJob) {
let job_id = job.id;
debug!(%job_id, "test_generation_worker: claimed job");
match process_job(state, job).await {
Ok(result) => {
match TestGenerationJob::complete_success(state.db.pool(), job_id, &result).await {
Ok(true) => {}
Ok(false) => {
debug!(%job_id, "test_generation_worker: success write was no-op (likely cancelled)");
}
Err(e) => {
error!(%job_id, error = ?e, "test_generation_worker: success write failed");
}
}
}
Err(reason) => {
let reason_str = reason.to_string();
warn!(%job_id, error = %reason_str, "test_generation_worker: job failed");
let _ =
TestGenerationJob::complete_failure(state.db.pool(), job_id, reason_str.as_str())
.await;
}
}
}
async fn process_job(state: &AppState, job: TestGenerationJob) -> Result<Value, WorkerError> {
let org = Organization::find_by_id(state.db.pool(), job.org_id)
.await
.map_err(|e| WorkerError::Internal(format!("org lookup failed: {e}")))?
.ok_or_else(|| WorkerError::Internal("Organization missing for job".into()))?;
let is_paid_plan = matches!(org.plan(), Plan::Pro | Plan::Team);
let byok = load_byok_config(state, job.org_id).await?;
let provider = pick_provider(is_paid_plan, byok);
let selection = provider.selection();
let quota = check_ai_quota(state, &org, selection)
.await
.map_err(|e| WorkerError::Internal(format!("quota check failed: {e:?}")))?;
if !quota.allowed {
return Err(WorkerError::ProviderUnavailable(
quota.deny_reason.unwrap_or_else(|| "AI quota exceeded".into()),
));
}
let filter = parse_filter(&job.captures_filter);
let captures = fetch_captures(state.db.pool(), job.workspace_id, &filter)
.await
.map_err(|e| WorkerError::Internal(format!("capture sampling failed: {e}")))?;
if captures.is_empty() {
return Err(WorkerError::EmptyCorpus(
"No matching captures found in this workspace. Record some traffic first or relax the filter.".into(),
));
}
let (call, provider_label) = build_call_for_provider(&provider, &job.prompt, &captures)?;
let llm_result = call_llm(call).await.map_err(|e| WorkerError::LlmCall(format!("{e:?}")))?;
let total_tokens = llm_result.total_tokens() as i64;
if let Err(e) = record_ai_usage(state, org.id, selection, total_tokens).await {
warn!(
job_id = %job.id,
org_id = %org.id,
tokens = total_tokens,
error = ?e,
"test_generation_worker: usage metering failed",
);
}
state
.store
.record_audit_event(
org.id,
job.created_by,
AuditEventType::AiUsage,
format!("AI test-generation completion via {} provider", provider_label),
Some(json!({
"handler": "test_generation_worker.process_job",
"job_id": job.id,
"provider": provider_label,
"prompt_tokens": llm_result.prompt_tokens,
"completion_tokens": llm_result.completion_tokens,
"total_tokens": llm_result.total_tokens(),
})),
None,
None,
)
.await;
let scenarios = parse_scenarios(&llm_result.content);
Ok(build_result_value(
&llm_result,
&provider_label,
selection,
scenarios,
captures.len(),
))
}
fn build_call_for_provider(
provider: &Provider,
user_prompt: &str,
captures: &[CaptureSample],
) -> Result<(LlmCall, String), WorkerError> {
let (system, user) = build_prompt(user_prompt, captures);
match provider {
Provider::Disabled => Err(WorkerError::ProviderUnavailable(
"AI is not available — add a BYOK key or upgrade your plan".into(),
)),
Provider::Byok(cfg) => {
let api_key = decrypt_api_key(&cfg.api_key)
.map_err(|e| WorkerError::Internal(format!("BYOK key decrypt failed: {e:?}")))?;
let provider_label = cfg.provider.clone();
let call = LlmCall {
provider: cfg.provider.clone(),
model: cfg.model.clone().unwrap_or_else(|| "gpt-4o-mini".into()),
api_key,
base_url: cfg.base_url.clone(),
system,
user,
temperature: 0.2,
max_tokens: LLM_MAX_COMPLETION_TOKENS,
};
Ok((call, provider_label))
}
Provider::Platform => {
let api_key = std::env::var("MOCKFORGE_PLATFORM_LLM_API_KEY")
.map_err(|_| WorkerError::ProviderUnavailable(
"Platform LLM not configured — set MOCKFORGE_PLATFORM_LLM_API_KEY on the registry or add a BYOK key.".into(),
))?;
let provider_name = std::env::var("MOCKFORGE_PLATFORM_LLM_PROVIDER")
.unwrap_or_else(|_| "openai".into());
let model = std::env::var("MOCKFORGE_PLATFORM_LLM_MODEL")
.unwrap_or_else(|_| "gpt-4o-mini".into());
let base_url = std::env::var("MOCKFORGE_PLATFORM_LLM_ENDPOINT").ok();
let provider_label = provider_name.clone();
let call = LlmCall {
provider: provider_name,
model,
api_key,
base_url,
system,
user,
temperature: 0.2,
max_tokens: LLM_MAX_COMPLETION_TOKENS,
};
Ok((call, provider_label))
}
}
}
#[derive(Debug, Default)]
struct CaptureFilter {
method: Option<String>,
path_contains: Option<String>,
status_min: Option<i32>,
status_max: Option<i32>,
limit: i64,
}
fn parse_filter(raw: &Value) -> CaptureFilter {
let mut f = CaptureFilter {
limit: MAX_CAPTURES,
..CaptureFilter::default()
};
let Some(obj) = raw.as_object() else {
return f;
};
if let Some(s) = obj.get("method").and_then(|v| v.as_str()) {
f.method = Some(s.to_uppercase());
}
if let Some(s) = obj.get("path_contains").and_then(|v| v.as_str()) {
f.path_contains = Some(s.to_string());
}
if let Some(n) = obj.get("status_min").and_then(|v| v.as_i64()) {
f.status_min = Some(n.clamp(100, 599) as i32);
}
if let Some(n) = obj.get("status_max").and_then(|v| v.as_i64()) {
f.status_max = Some(n.clamp(100, 599) as i32);
}
if let Some(n) = obj.get("limit").and_then(|v| v.as_i64()) {
f.limit = n.clamp(1, MAX_CAPTURES);
}
f
}
#[derive(Debug, Clone, FromRow)]
struct CaptureSample {
method: String,
path: String,
#[sqlx(rename = "effective_status")]
status: i32,
duration_ms: i32,
}
async fn fetch_captures(
pool: &PgPool,
workspace_id: Uuid,
filter: &CaptureFilter,
) -> sqlx::Result<Vec<CaptureSample>> {
sqlx::query_as::<_, CaptureSample>(
r#"
SELECT method, path,
COALESCE(response_status_code, status_code, 0) AS effective_status,
COALESCE(duration_ms, 0) AS duration_ms
FROM runtime_captures
WHERE workspace_id = $1
AND ($2::text IS NULL OR UPPER(method) = $2)
AND ($3::text IS NULL OR position($3 IN path) > 0)
AND ($4::int IS NULL
OR COALESCE(response_status_code, status_code, 0) >= $4)
AND ($5::int IS NULL
OR COALESCE(response_status_code, status_code, 0) <= $5)
ORDER BY occurred_at DESC
LIMIT $6
"#,
)
.bind(workspace_id)
.bind(filter.method.as_deref())
.bind(filter.path_contains.as_deref())
.bind(filter.status_min)
.bind(filter.status_max)
.bind(filter.limit)
.fetch_all(pool)
.await
}
fn build_prompt(user_prompt: &str, captures: &[CaptureSample]) -> (String, String) {
let system = "You are a senior test engineer. Given a sample of recent API requests, you propose concise test scenarios that would catch realistic regressions. Output ONLY a single JSON object on one line of the form: {\"scenarios\": [{\"name\": \"...\", \"description\": \"...\", \"method\": \"GET\", \"path\": \"/foo\", \"expected_status\": 200, \"rationale\": \"...\"}, ...]}. No prose, no code fences, no explanation. Up to 10 scenarios.".to_string();
let mut lines = String::with_capacity(64 * captures.len());
lines.push_str("Recent captures (method | path | status | duration_ms):\n");
for c in captures {
lines.push_str(&format!("{} {} {} {}\n", c.method, c.path, c.status, c.duration_ms));
}
let extra = if user_prompt.trim().is_empty() {
String::new()
} else {
format!("\n\nFocus area from the user:\n{}", user_prompt.trim())
};
let user = format!("{lines}{extra}");
(system, user)
}
fn parse_scenarios(content: &str) -> Option<Value> {
let trimmed = content.trim();
let cleaned = trimmed
.trim_start_matches("```json")
.trim_start_matches("```")
.trim_end_matches("```")
.trim();
if let Ok(v) = serde_json::from_str::<Value>(cleaned) {
return Some(v);
}
if let (Some(start), Some(end)) = (cleaned.find('{'), cleaned.rfind('}')) {
if end > start {
if let Ok(v) = serde_json::from_str::<Value>(&cleaned[start..=end]) {
return Some(v);
}
}
}
None
}
fn build_result_value(
llm: &LlmResult,
provider_label: &str,
selection: ProviderSelection,
parsed: Option<Value>,
captures_sampled: usize,
) -> Value {
let billing = match selection {
ProviderSelection::Byok => "byok",
ProviderSelection::Platform => "platform",
ProviderSelection::Disabled => "disabled",
};
json!({
"scenarios": parsed.as_ref().and_then(|v| v.get("scenarios").cloned()),
"raw_parsed": parsed,
"raw_content": llm.content,
"model_meta": {
"provider": provider_label,
"billing": billing,
"prompt_tokens": llm.prompt_tokens,
"completion_tokens": llm.completion_tokens,
},
"captures_sampled": captures_sampled,
})
}
async fn load_byok_config(
state: &AppState,
org_id: Uuid,
) -> Result<Option<BYOKConfig>, WorkerError> {
let setting = state
.store
.get_org_setting(org_id, "byok")
.await
.map_err(|e| WorkerError::Internal(format!("byok lookup failed: {e:?}")))?;
let Some(setting) = setting else {
return Ok(None);
};
let cfg: BYOKConfig = match serde_json::from_value(setting.setting_value) {
Ok(c) => c,
Err(_) => return Ok(None),
};
if !cfg.enabled || cfg.api_key.is_empty() {
return Ok(None);
}
Ok(Some(cfg))
}
#[derive(Debug)]
enum WorkerError {
ProviderUnavailable(String),
EmptyCorpus(String),
LlmCall(String),
Internal(String),
}
impl std::fmt::Display for WorkerError {
fn fmt(&self, f: &mut std::fmt::Formatter<'_>) -> std::fmt::Result {
match self {
Self::ProviderUnavailable(m)
| Self::EmptyCorpus(m)
| Self::LlmCall(m)
| Self::Internal(m) => write!(f, "{m}"),
}
}
}
#[cfg(test)]
#[derive(Debug, serde::Deserialize)]
struct FilterRoundTrip {
#[serde(default)]
method: Option<String>,
#[serde(default)]
path_contains: Option<String>,
#[serde(default)]
status_min: Option<i32>,
#[serde(default)]
status_max: Option<i32>,
#[serde(default)]
limit: Option<i64>,
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn parse_filter_empty_defaults_to_max_captures() {
let f = parse_filter(&json!({}));
assert!(f.method.is_none());
assert!(f.path_contains.is_none());
assert!(f.status_min.is_none());
assert!(f.status_max.is_none());
assert_eq!(f.limit, MAX_CAPTURES);
}
#[test]
fn parse_filter_normalises_method_uppercase_and_clamps_limit() {
let f = parse_filter(&json!({
"method": "post",
"limit": 1_000_000,
}));
assert_eq!(f.method.as_deref(), Some("POST"));
assert_eq!(f.limit, MAX_CAPTURES, "limit clamped to MAX_CAPTURES");
}
#[test]
fn parse_filter_clamps_status_to_http_range() {
let f = parse_filter(&json!({
"status_min": 50,
"status_max": 999,
}));
assert_eq!(f.status_min, Some(100));
assert_eq!(f.status_max, Some(599));
}
#[test]
fn parse_filter_round_trips_path_contains() {
let f = parse_filter(&json!({"path_contains": "/auth/"}));
assert_eq!(f.path_contains.as_deref(), Some("/auth/"));
}
#[test]
fn parse_filter_handles_non_object_input() {
let f = parse_filter(&Value::Null);
assert!(f.method.is_none());
let f = parse_filter(&json!("not an object"));
assert!(f.method.is_none());
}
#[test]
fn build_prompt_includes_captures_and_user_focus() {
let captures = vec![
CaptureSample {
method: "GET".into(),
path: "/users".into(),
status: 200,
duration_ms: 12,
},
CaptureSample {
method: "POST".into(),
path: "/users".into(),
status: 201,
duration_ms: 34,
},
];
let (system, user) = build_prompt("focus on auth", &captures);
assert!(system.contains("scenarios"));
assert!(user.contains("GET /users 200 12"));
assert!(user.contains("POST /users 201 34"));
assert!(user.contains("focus on auth"));
}
#[test]
fn build_prompt_omits_user_focus_block_when_empty() {
let captures = vec![CaptureSample {
method: "GET".into(),
path: "/".into(),
status: 200,
duration_ms: 1,
}];
let (_, user) = build_prompt(" ", &captures);
assert!(!user.contains("Focus area"));
}
#[test]
fn parse_scenarios_accepts_plain_json() {
let v = parse_scenarios(r#"{"scenarios": [{"name": "ok"}]}"#).unwrap();
assert_eq!(v["scenarios"][0]["name"], "ok");
}
#[test]
fn parse_scenarios_strips_code_fences() {
let raw = "```json\n{\"scenarios\": [1, 2]}\n```";
let v = parse_scenarios(raw).unwrap();
assert_eq!(v["scenarios"].as_array().unwrap().len(), 2);
}
#[test]
fn parse_scenarios_recovers_from_prose_wrapper() {
let raw =
"Sure! Here you go:\n{\"scenarios\": [\"a\", \"b\"]}\nLet me know if you need more.";
let v = parse_scenarios(raw).unwrap();
assert_eq!(v["scenarios"].as_array().unwrap().len(), 2);
}
#[test]
fn parse_scenarios_returns_none_on_unrecoverable_input() {
assert!(parse_scenarios("not json at all").is_none());
}
#[test]
fn build_result_value_preserves_raw_content_when_parse_fails() {
let llm = LlmResult {
content: "garbage".into(),
prompt_tokens: 100,
completion_tokens: 50,
};
let v = build_result_value(&llm, "openai", ProviderSelection::Byok, None, 5);
assert_eq!(v["raw_content"], "garbage");
assert!(v["scenarios"].is_null());
assert!(v["raw_parsed"].is_null());
assert_eq!(v["model_meta"]["prompt_tokens"], 100);
assert_eq!(v["model_meta"]["billing"], "byok");
assert_eq!(v["captures_sampled"], 5);
}
#[test]
fn build_result_value_hoists_scenarios_field_and_tags_platform_billing() {
let llm = LlmResult {
content: r#"{"scenarios": [{"name": "happy"}]}"#.into(),
prompt_tokens: 0,
completion_tokens: 0,
};
let parsed = parse_scenarios(&llm.content);
let v = build_result_value(&llm, "openai", ProviderSelection::Platform, parsed, 3);
assert_eq!(v["scenarios"][0]["name"], "happy");
assert_eq!(v["model_meta"]["billing"], "platform");
assert_eq!(v["captures_sampled"], 3);
}
#[test]
fn build_call_for_provider_disabled_returns_clear_error() {
let captures = vec![CaptureSample {
method: "GET".into(),
path: "/".into(),
status: 200,
duration_ms: 1,
}];
let err = build_call_for_provider(&Provider::Disabled, "", &captures).unwrap_err();
match err {
WorkerError::ProviderUnavailable(msg) => {
assert!(msg.contains("BYOK") || msg.contains("upgrade"));
}
other => panic!("expected ProviderUnavailable, got {other:?}"),
}
}
#[test]
fn build_call_for_provider_platform_requires_env() {
let prev = std::env::var("MOCKFORGE_PLATFORM_LLM_API_KEY").ok();
std::env::remove_var("MOCKFORGE_PLATFORM_LLM_API_KEY");
let captures = vec![CaptureSample {
method: "GET".into(),
path: "/".into(),
status: 200,
duration_ms: 1,
}];
let err = build_call_for_provider(&Provider::Platform, "", &captures).unwrap_err();
match err {
WorkerError::ProviderUnavailable(msg) => {
assert!(
msg.contains("MOCKFORGE_PLATFORM_LLM_API_KEY"),
"expected env-var hint in message: {msg}"
);
}
other => panic!("expected ProviderUnavailable, got {other:?}"),
}
if let Some(v) = prev {
std::env::set_var("MOCKFORGE_PLATFORM_LLM_API_KEY", v);
}
}
#[test]
fn worker_error_display_unwraps_message() {
assert_eq!(format!("{}", WorkerError::EmptyCorpus("no rows".into())), "no rows");
}
#[test]
fn filter_round_trip_smoke() {
let f: FilterRoundTrip = serde_json::from_value(json!({
"method": "GET",
"status_min": 400,
"limit": 10,
}))
.unwrap();
assert_eq!(f.method.as_deref(), Some("GET"));
assert_eq!(f.status_min, Some(400));
assert_eq!(f.limit, Some(10));
assert!(f.path_contains.is_none());
}
}