Skip to main content

tatara_testing/
server.rs

1//! In-process HTTP test server for tatara API testing.
2//!
3//! Provides a fully functional tatara REST API backed by an InMemoryStore,
4//! without requiring Raft consensus, gossip, or any networking infrastructure.
5//!
6//! The test server uses Tower's `oneshot` pattern for zero-overhead HTTP testing
7//! (no TCP sockets, no port allocation).
8
9use 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
26/// In-process test server for tatara API integration testing.
27///
28/// Uses Tower's `oneshot` pattern — no TCP sockets, no port allocation.
29/// Each request is processed directly through the axum Router.
30pub 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    /// Create a new test server with an empty store.
42    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            // Jobs
51            .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            // Allocations
60            .route("/api/v1/allocations", get(list_allocations))
61            .route("/api/v1/allocations/{alloc_id}", get(get_allocation))
62            // Nodes
63            .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            // Events
70            .route("/api/v1/events", get(list_events))
71            // Releases
72            .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    /// Create a test server with a pre-populated store.
88    pub fn with_store(store: Arc<InMemoryStore>) -> Self {
89        let state = TestState {
90            store: store.clone(),
91        };
92
93        // Build same router
94        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    /// Send a GET request to the test server.
129    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    /// Send a POST request with a JSON body.
138    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
160/// Response from the test server with convenient assertion methods.
161pub 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    /// Assert the response status is 200 OK.
177    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    /// Assert the response status matches.
188    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    /// Deserialize the response body as JSON.
200    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    /// Get the response body as a string.
212    pub fn body_text(&self) -> String {
213        String::from_utf8_lossy(&self.body).to_string()
214    }
215}
216
217// ── Handlers (mirror production API but use InMemoryStore) ──
218
219async 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    // Supersede other active releases
448    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}