1use axum::body::Body;
10use axum::extract::{Path, Query, State};
11use axum::http::{Request, StatusCode};
12use axum::routing::{get, post};
13use axum::{Json, Router};
14use serde::Deserialize;
15use std::sync::Arc;
16use tower::ServiceExt;
17use uuid::Uuid;
18
19use super::store::InMemoryStore;
20use tatara_core::cluster::types::{JobVersionEntry, NodeMeta};
21use tatara_core::domain::allocation::Allocation;
22use tatara_core::domain::event::EventKind;
23use tatara_core::domain::job::{Job, JobSpec, JobStatus};
24use tatara_core::domain::release::{CreateReleaseRequest, Release, ReleaseStatus};
25
26pub struct TestServer {
31 pub store: Arc<InMemoryStore>,
32 router: Router,
33}
34
35#[derive(Clone)]
36struct TestState {
37 store: Arc<InMemoryStore>,
38}
39
40impl TestServer {
41 pub fn new() -> Self {
43 let store = Arc::new(InMemoryStore::new());
44 let state = TestState {
45 store: store.clone(),
46 };
47
48 let router = Router::new()
49 .route("/health", get(health))
50 .route("/api/v1/jobs", get(list_jobs).post(submit_job))
52 .route("/api/v1/jobs/{job_id}", get(get_job))
53 .route("/api/v1/jobs/{job_id}/stop", post(stop_job))
54 .route("/api/v1/jobs/{job_id}/history", get(get_job_history))
55 .route(
56 "/api/v1/jobs/{job_id}/rollback/{version}",
57 post(rollback_job),
58 )
59 .route("/api/v1/allocations", get(list_allocations))
61 .route("/api/v1/allocations/{alloc_id}", get(get_allocation))
62 .route("/api/v1/nodes", get(list_nodes))
64 .route("/api/v1/nodes/{node_id}/drain", post(drain_node))
65 .route(
66 "/api/v1/nodes/{node_id}/eligibility",
67 post(set_node_eligibility),
68 )
69 .route("/api/v1/events", get(list_events))
71 .route("/api/v1/releases", get(list_releases).post(create_release))
73 .route("/api/v1/releases/{release_id}", get(get_release))
74 .route(
75 "/api/v1/releases/{release_id}/promote",
76 post(promote_release),
77 )
78 .route(
79 "/api/v1/releases/{release_id}/rollback",
80 post(rollback_release),
81 )
82 .with_state(state);
83
84 Self { store, router }
85 }
86
87 pub fn with_store(store: Arc<InMemoryStore>) -> Self {
89 let state = TestState {
90 store: store.clone(),
91 };
92
93 let router = Router::new()
95 .route("/health", get(health))
96 .route("/api/v1/jobs", get(list_jobs).post(submit_job))
97 .route("/api/v1/jobs/{job_id}", get(get_job))
98 .route("/api/v1/jobs/{job_id}/stop", post(stop_job))
99 .route("/api/v1/jobs/{job_id}/history", get(get_job_history))
100 .route(
101 "/api/v1/jobs/{job_id}/rollback/{version}",
102 post(rollback_job),
103 )
104 .route("/api/v1/allocations", get(list_allocations))
105 .route("/api/v1/allocations/{alloc_id}", get(get_allocation))
106 .route("/api/v1/nodes", get(list_nodes))
107 .route("/api/v1/nodes/{node_id}/drain", post(drain_node))
108 .route(
109 "/api/v1/nodes/{node_id}/eligibility",
110 post(set_node_eligibility),
111 )
112 .route("/api/v1/events", get(list_events))
113 .route("/api/v1/releases", get(list_releases).post(create_release))
114 .route("/api/v1/releases/{release_id}", get(get_release))
115 .route(
116 "/api/v1/releases/{release_id}/promote",
117 post(promote_release),
118 )
119 .route(
120 "/api/v1/releases/{release_id}/rollback",
121 post(rollback_release),
122 )
123 .with_state(state);
124
125 Self { store, router }
126 }
127
128 pub async fn get(&self, uri: &str) -> TestResponse {
130 let request = Request::builder().uri(uri).body(Body::empty()).unwrap();
131
132 let response = self.router.clone().oneshot(request).await.unwrap();
133
134 TestResponse::from_response(response).await
135 }
136
137 pub async fn post<T: serde::Serialize>(&self, uri: &str, body: &T) -> TestResponse {
139 let body_bytes = serde_json::to_vec(body).unwrap();
140
141 let request = Request::builder()
142 .method("POST")
143 .uri(uri)
144 .header("content-type", "application/json")
145 .body(Body::from(body_bytes))
146 .unwrap();
147
148 let response = self.router.clone().oneshot(request).await.unwrap();
149
150 TestResponse::from_response(response).await
151 }
152}
153
154impl Default for TestServer {
155 fn default() -> Self {
156 Self::new()
157 }
158}
159
160pub struct TestResponse {
162 pub status: StatusCode,
163 pub body: Vec<u8>,
164}
165
166impl TestResponse {
167 async fn from_response(response: axum::http::Response<Body>) -> Self {
168 let status = response.status();
169 let body = axum::body::to_bytes(response.into_body(), usize::MAX)
170 .await
171 .unwrap()
172 .to_vec();
173 Self { status, body }
174 }
175
176 pub fn assert_ok(&self) {
178 assert_eq!(
179 self.status,
180 StatusCode::OK,
181 "Expected 200 OK, got {}. Body: {}",
182 self.status,
183 self.body_text()
184 );
185 }
186
187 pub fn assert_status(&self, expected: StatusCode) {
189 assert_eq!(
190 self.status,
191 expected,
192 "Expected {}, got {}. Body: {}",
193 expected,
194 self.status,
195 self.body_text()
196 );
197 }
198
199 pub fn json<T: serde::de::DeserializeOwned>(&self) -> T {
201 serde_json::from_slice(&self.body).unwrap_or_else(|e| {
202 panic!(
203 "Failed to deserialize response body as {}: {}. Body: {}",
204 std::any::type_name::<T>(),
205 e,
206 self.body_text()
207 )
208 })
209 }
210
211 pub fn body_text(&self) -> String {
213 String::from_utf8_lossy(&self.body).to_string()
214 }
215}
216
217async fn health() -> &'static str {
220 "ok"
221}
222
223async fn submit_job(
224 State(state): State<TestState>,
225 Json(spec): Json<JobSpec>,
226) -> Result<Json<Job>, (StatusCode, String)> {
227 let job = spec.into_job();
228 let job = state
229 .store
230 .put_job(job)
231 .await
232 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
233 Ok(Json(job))
234}
235
236async fn list_jobs(State(state): State<TestState>) -> Json<Vec<Job>> {
237 Json(state.store.list_jobs().await)
238}
239
240#[derive(serde::Serialize)]
241struct JobDetail {
242 job: Job,
243 allocations: Vec<Allocation>,
244}
245
246async fn get_job(
247 State(state): State<TestState>,
248 Path(job_id): Path<String>,
249) -> Result<Json<JobDetail>, (StatusCode, String)> {
250 let job = state
251 .store
252 .get_job(&job_id)
253 .await
254 .ok_or((StatusCode::NOT_FOUND, "Job not found".to_string()))?;
255
256 let allocations = state.store.list_allocations_for_job(&job_id).await;
257 Ok(Json(JobDetail { job, allocations }))
258}
259
260async fn stop_job(
261 State(state): State<TestState>,
262 Path(job_id): Path<String>,
263) -> Result<Json<Job>, (StatusCode, String)> {
264 let job = state
265 .store
266 .update_job_status(&job_id, JobStatus::Dead)
267 .await
268 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
269 Ok(Json(job))
270}
271
272async fn get_job_history(
273 State(state): State<TestState>,
274 Path(job_id): Path<String>,
275) -> Result<Json<Vec<JobVersionEntry>>, (StatusCode, String)> {
276 let history = state.store.get_job_history(&job_id).await;
277 if history.is_empty() {
278 if state.store.get_job(&job_id).await.is_none() {
279 return Err((StatusCode::NOT_FOUND, "Job not found".to_string()));
280 }
281 }
282 Ok(Json(history))
283}
284
285async fn rollback_job(
286 State(state): State<TestState>,
287 Path((job_id, version)): Path<(String, u64)>,
288) -> Result<Json<Job>, (StatusCode, String)> {
289 let job = state
290 .store
291 .rollback_job(&job_id, version)
292 .await
293 .map_err(|e| {
294 let msg = e.to_string();
295 if msg.contains("not found") {
296 (StatusCode::NOT_FOUND, msg)
297 } else {
298 (StatusCode::INTERNAL_SERVER_ERROR, msg)
299 }
300 })?;
301 Ok(Json(job))
302}
303
304async fn list_allocations(State(state): State<TestState>) -> Json<Vec<Allocation>> {
305 Json(state.store.list_allocations().await)
306}
307
308async fn get_allocation(
309 State(state): State<TestState>,
310 Path(alloc_id): Path<String>,
311) -> Result<Json<Allocation>, (StatusCode, String)> {
312 let id: Uuid = alloc_id
313 .parse()
314 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid allocation ID".to_string()))?;
315
316 state
317 .store
318 .get_allocation(&id)
319 .await
320 .map(Json)
321 .ok_or((StatusCode::NOT_FOUND, "Allocation not found".to_string()))
322}
323
324async fn list_nodes(State(state): State<TestState>) -> Json<Vec<NodeMeta>> {
325 Json(state.store.list_nodes().await)
326}
327
328#[derive(Deserialize)]
329struct DrainRequest {
330 #[serde(default)]
331 deadline_secs: Option<u64>,
332}
333
334async fn drain_node(
335 State(state): State<TestState>,
336 Path(node_id): Path<String>,
337 Json(body): Json<DrainRequest>,
338) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
339 let id: u64 = node_id
340 .parse()
341 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid node ID".to_string()))?;
342
343 state
344 .store
345 .drain_node(id)
346 .await
347 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
348
349 Ok(Json(serde_json::json!({
350 "node_id": id,
351 "status": "draining",
352 "deadline_secs": body.deadline_secs,
353 })))
354}
355
356#[derive(Deserialize)]
357struct EligibilityRequest {
358 eligible: bool,
359}
360
361async fn set_node_eligibility(
362 State(state): State<TestState>,
363 Path(node_id): Path<String>,
364 Json(body): Json<EligibilityRequest>,
365) -> Result<Json<serde_json::Value>, (StatusCode, String)> {
366 let id: u64 = node_id
367 .parse()
368 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid node ID".to_string()))?;
369
370 state
371 .store
372 .set_node_eligibility(id, body.eligible)
373 .await
374 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
375
376 Ok(Json(serde_json::json!({
377 "node_id": id,
378 "eligible": body.eligible,
379 })))
380}
381
382#[derive(Deserialize)]
383struct EventQuery {
384 kind: Option<String>,
385 since: Option<String>,
386}
387
388async fn list_events(
389 State(state): State<TestState>,
390 params: Query<EventQuery>,
391) -> Json<Vec<tatara_core::domain::event::Event>> {
392 let kind = params.kind.as_deref().and_then(EventKind::from_str_opt);
393
394 let since = params
395 .since
396 .as_deref()
397 .and_then(|s| chrono::DateTime::parse_from_rfc3339(s).ok())
398 .map(|dt| dt.with_timezone(&chrono::Utc));
399
400 Json(state.store.list_events(kind.as_ref(), since).await)
401}
402
403async fn list_releases(State(state): State<TestState>) -> Json<Vec<Release>> {
404 Json(state.store.list_releases().await)
405}
406
407async fn create_release(
408 State(state): State<TestState>,
409 Json(req): Json<CreateReleaseRequest>,
410) -> Result<Json<Release>, (StatusCode, String)> {
411 let mut release = Release::new(req.name, req.flake_ref, req.job_id);
412 release.flake_rev = req.flake_rev;
413 release.status = ReleaseStatus::Active;
414
415 let release = state
416 .store
417 .put_release(release)
418 .await
419 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
420 Ok(Json(release))
421}
422
423async fn get_release(
424 State(state): State<TestState>,
425 Path(release_id): Path<String>,
426) -> Result<Json<Release>, (StatusCode, String)> {
427 let id: Uuid = release_id
428 .parse()
429 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid release ID".to_string()))?;
430
431 state
432 .store
433 .get_release(&id)
434 .await
435 .map(Json)
436 .ok_or((StatusCode::NOT_FOUND, "Release not found".to_string()))
437}
438
439async fn promote_release(
440 State(state): State<TestState>,
441 Path(release_id): Path<String>,
442) -> Result<Json<Release>, (StatusCode, String)> {
443 let id: Uuid = release_id
444 .parse()
445 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid release ID".to_string()))?;
446
447 let releases = state.store.list_releases().await;
449 for rel in &releases {
450 if rel.id != id && rel.status == ReleaseStatus::Active {
451 let _ = state
452 .store
453 .update_release_status(rel.id, ReleaseStatus::Superseded)
454 .await;
455 }
456 }
457
458 let release = state
459 .store
460 .update_release_status(id, ReleaseStatus::Active)
461 .await
462 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
463 Ok(Json(release))
464}
465
466async fn rollback_release(
467 State(state): State<TestState>,
468 Path(release_id): Path<String>,
469) -> Result<Json<Release>, (StatusCode, String)> {
470 let id: Uuid = release_id
471 .parse()
472 .map_err(|_| (StatusCode::BAD_REQUEST, "Invalid release ID".to_string()))?;
473
474 let release = state
475 .store
476 .update_release_status(id, ReleaseStatus::RolledBack)
477 .await
478 .map_err(|e| (StatusCode::INTERNAL_SERVER_ERROR, e.to_string()))?;
479 Ok(Json(release))
480}