use std::convert::Infallible;
use std::time::Duration;
use axum::{
extract::{Path, State},
http::HeaderMap,
response::sse::{Event, KeepAlive, Sse},
Json,
};
use futures_util::stream::{self, Stream};
use serde::Deserialize;
use serde_json::{json, Value};
use sqlx::PgPool;
use uuid::Uuid;
use crate::{
error::{ApiError, ApiResult},
middleware::{resolve_org_context, AuthUser},
AppState,
};
use mockforge_registry_core::models::{
test_generation_job::{CreateTestGenerationJob, TestGenerationJob},
CloudWorkspace,
};
const LIST_LIMIT: i64 = 100;
const MAX_PROMPT_BYTES: usize = 8 * 1024;
#[derive(Debug, Deserialize)]
pub struct CreateJobRequest {
#[serde(default)]
pub prompt: String,
#[serde(default)]
pub captures_filter: Value,
}
const MAX_FILTER_BYTES: usize = 16 * 1024;
pub async fn create_job(
State(state): State<AppState>,
AuthUser(user_id): AuthUser,
Path(workspace_id): Path<Uuid>,
headers: HeaderMap,
Json(body): Json<CreateJobRequest>,
) -> ApiResult<Json<TestGenerationJob>> {
let workspace = authorize_workspace(&state, user_id, &headers, workspace_id).await?;
if body.prompt.len() > MAX_PROMPT_BYTES {
return Err(ApiError::InvalidRequest(format!(
"prompt exceeds {MAX_PROMPT_BYTES} byte limit"
)));
}
let filter_size = serde_json::to_vec(&body.captures_filter).map(|v| v.len()).unwrap_or(0);
if filter_size > MAX_FILTER_BYTES {
return Err(ApiError::InvalidRequest(format!(
"captures_filter exceeds {MAX_FILTER_BYTES} byte limit"
)));
}
if !body.captures_filter.is_object() && !body.captures_filter.is_null() {
return Err(ApiError::InvalidRequest("captures_filter must be a JSON object".into()));
}
let captures_filter = if body.captures_filter.is_null() {
json!({})
} else {
body.captures_filter
};
let row = TestGenerationJob::create(
state.db.pool(),
CreateTestGenerationJob {
workspace_id: workspace.id,
org_id: workspace.org_id,
prompt: &body.prompt,
captures_filter: &captures_filter,
created_by: Some(user_id),
},
)
.await?;
tracing::info!(
job_id = %row.id,
workspace_id = %workspace.id,
org_id = %workspace.org_id,
prompt_len = body.prompt.len(),
"test-generation job queued"
);
Ok(Json(row))
}
pub async fn list_jobs(
State(state): State<AppState>,
AuthUser(user_id): AuthUser,
Path(workspace_id): Path<Uuid>,
headers: HeaderMap,
) -> ApiResult<Json<Vec<TestGenerationJob>>> {
let workspace = authorize_workspace(&state, user_id, &headers, workspace_id).await?;
let jobs =
TestGenerationJob::list_by_workspace(state.db.pool(), workspace.id, LIST_LIMIT).await?;
Ok(Json(jobs))
}
pub async fn get_job(
State(state): State<AppState>,
AuthUser(user_id): AuthUser,
Path((workspace_id, job_id)): Path<(Uuid, Uuid)>,
headers: HeaderMap,
) -> ApiResult<Json<TestGenerationJob>> {
let workspace = authorize_workspace(&state, user_id, &headers, workspace_id).await?;
let job = TestGenerationJob::find_in_workspace(state.db.pool(), workspace.id, job_id)
.await?
.ok_or_else(|| ApiError::InvalidRequest("Job not found".into()))?;
Ok(Json(job))
}
pub async fn cancel_job(
State(state): State<AppState>,
AuthUser(user_id): AuthUser,
Path((workspace_id, job_id)): Path<(Uuid, Uuid)>,
headers: HeaderMap,
) -> ApiResult<Json<Value>> {
let workspace = authorize_workspace(&state, user_id, &headers, workspace_id).await?;
let changed = TestGenerationJob::cancel(state.db.pool(), workspace.id, job_id).await?;
Ok(Json(json!({
"cancelled": changed,
})))
}
pub async fn stream_job(
State(state): State<AppState>,
AuthUser(user_id): AuthUser,
Path((workspace_id, job_id)): Path<(Uuid, Uuid)>,
headers: HeaderMap,
) -> ApiResult<Sse<impl Stream<Item = Result<Event, Infallible>>>> {
let workspace = authorize_workspace(&state, user_id, &headers, workspace_id).await?;
let cursor = JobStreamCursor {
pool: state.db.pool().clone(),
workspace_id: workspace.id,
job_id,
last_snapshot: None,
terminal_emitted: false,
};
let stream = stream::unfold(cursor, advance_job_stream);
Ok(Sse::new(stream).keep_alive(KeepAlive::default()))
}
#[derive(Debug, Clone, PartialEq, Eq)]
struct JobSnapshot {
status: String,
started_at_set: bool,
finished_at_set: bool,
has_result: bool,
has_error: bool,
}
impl JobSnapshot {
fn from_job(j: &TestGenerationJob) -> Self {
Self {
status: j.status.clone(),
started_at_set: j.started_at.is_some(),
finished_at_set: j.finished_at.is_some(),
has_result: j.result.is_some(),
has_error: j.error.is_some(),
}
}
fn is_terminal(&self) -> bool {
matches!(self.status.as_str(), "succeeded" | "failed" | "cancelled")
}
}
struct JobStreamCursor {
pool: PgPool,
workspace_id: Uuid,
job_id: Uuid,
last_snapshot: Option<JobSnapshot>,
terminal_emitted: bool,
}
async fn advance_job_stream(
mut cursor: JobStreamCursor,
) -> Option<(Result<Event, Infallible>, JobStreamCursor)> {
if cursor.terminal_emitted {
return None;
}
if cursor.last_snapshot.is_some() {
tokio::time::sleep(Duration::from_secs(1)).await;
}
let job = match TestGenerationJob::find_in_workspace(
&cursor.pool,
cursor.workspace_id,
cursor.job_id,
)
.await
{
Ok(Some(j)) => j,
Ok(None) => {
let evt = Event::default()
.event("not_found")
.data(json!({ "job_id": cursor.job_id }).to_string());
cursor.terminal_emitted = true;
return Some((Ok(evt), cursor));
}
Err(e) => {
let evt = Event::default()
.event("stream_error")
.data(json!({ "error": e.to_string() }).to_string());
cursor.terminal_emitted = true;
return Some((Ok(evt), cursor));
}
};
let snapshot = JobSnapshot::from_job(&job);
let unchanged = cursor.last_snapshot.as_ref().is_some_and(|s| s == &snapshot);
if unchanged {
let evt = Event::default().event("ping").data("{}");
return Some((Ok(evt), cursor));
}
let terminal = snapshot.is_terminal();
cursor.last_snapshot = Some(snapshot);
if terminal {
cursor.terminal_emitted = true;
}
let payload =
serde_json::to_value(&job).unwrap_or_else(|_| json!({ "error": "serialization failed" }));
let evt = Event::default().event("status_update").data(payload.to_string());
Some((Ok(evt), cursor))
}
async fn authorize_workspace(
state: &AppState,
user_id: Uuid,
headers: &HeaderMap,
workspace_id: Uuid,
) -> ApiResult<CloudWorkspace> {
let workspace = CloudWorkspace::find_by_id(state.db.pool(), workspace_id)
.await?
.ok_or_else(|| ApiError::InvalidRequest("Workspace not found".into()))?;
let ctx = resolve_org_context(state, user_id, headers, None)
.await
.map_err(|_| ApiError::InvalidRequest("Organization not found".into()))?;
if ctx.org_id != workspace.org_id {
return Err(ApiError::InvalidRequest("Workspace not found".into()));
}
Ok(workspace)
}
#[cfg(test)]
mod tests {
use super::*;
#[test]
fn create_request_defaults() {
let req: CreateJobRequest = serde_json::from_str("{}").unwrap();
assert_eq!(req.prompt, "");
assert!(req.captures_filter.is_null());
}
#[test]
fn create_request_accepts_prompt_and_filter() {
let req: CreateJobRequest =
serde_json::from_str(r#"{"prompt":"gen tests","captures_filter":{"status":">=400"}}"#)
.unwrap();
assert_eq!(req.prompt, "gen tests");
assert_eq!(req.captures_filter["status"], ">=400");
}
#[test]
fn prompt_length_cap_is_8kb() {
assert_eq!(MAX_PROMPT_BYTES, 8 * 1024);
}
#[test]
fn filter_size_cap_is_16kb() {
assert_eq!(MAX_FILTER_BYTES, 16 * 1024);
}
#[test]
fn list_limit_is_capped_at_100() {
assert_eq!(LIST_LIMIT, 100);
}
}