Skip to main content

spikard_codegen/sql/
openapi.rs

1//! Emit an OpenAPI 3.1 document from a slice of [`SqlRoute`].
2//!
3//! The spec is built as a raw `serde_json::Value` rather than reusing
4//! `crate::openapi::OpenApiSpec` because the existing struct is a subset that
5//! lacks several 3.1 idioms we need (array-typed `type`, `oneOf` for
6//! nullability, `enum`). Emitting as `Value` keeps this module decoupled and
7//! the output round-trips through any OpenAPI 3.1 consumer.
8
9use indexmap::IndexMap;
10use serde::{Deserialize, Serialize};
11use serde_json::{Map, Value, json};
12
13use super::annotations::{ApiKeyLocation, AuthRequirement, HttpMethod, HttpParamBinding};
14use super::route::SqlRoute;
15
16#[derive(Debug, Clone, Serialize, Deserialize)]
17pub struct OpenApiInfo {
18    pub title: String,
19    pub version: String,
20    #[serde(skip_serializing_if = "Option::is_none")]
21    pub description: Option<String>,
22}
23
24impl OpenApiInfo {
25    pub fn new(title: impl Into<String>, version: impl Into<String>) -> Self {
26        Self {
27            title: title.into(),
28            version: version.into(),
29            description: None,
30        }
31    }
32}
33
34/// Build an OpenAPI 3.1 document from a list of SQL-derived routes. The
35/// returned `Value` is ready to be `serde_json::to_writer_pretty`-ed to disk.
36pub fn openapi_from_routes(routes: &[SqlRoute], info: &OpenApiInfo) -> Value {
37    let (security_schemes, scheme_names) = collect_security_schemes(routes);
38
39    let mut paths: IndexMap<String, Map<String, Value>> = IndexMap::new();
40    for route in routes {
41        let entry = paths.entry(route.http.path.clone()).or_default();
42        let operation = build_operation(route, &scheme_names);
43        entry.insert(method_key(route.http.method).to_string(), operation);
44    }
45
46    let mut paths_obj = Map::new();
47    for (path, methods) in paths {
48        paths_obj.insert(path, Value::Object(methods));
49    }
50
51    let mut spec = Map::new();
52    spec.insert("openapi".into(), json!("3.1.0"));
53    spec.insert("info".into(), serde_json::to_value(info).expect("info serializes"));
54    spec.insert("paths".into(), Value::Object(paths_obj));
55
56    let mut components = Map::new();
57    if !security_schemes.is_empty() {
58        components.insert("securitySchemes".into(), Value::Object(security_schemes));
59    }
60    if !components.is_empty() {
61        spec.insert("components".into(), Value::Object(components));
62    }
63
64    Value::Object(spec)
65}
66
67fn build_operation(route: &SqlRoute, scheme_names: &std::collections::BTreeMap<AuthRequirement, String>) -> Value {
68    let mut op = Map::new();
69    op.insert("operationId".into(), json!(&route.operation_id));
70
71    if let Some(s) = &route.http.summary {
72        op.insert("summary".into(), json!(s));
73    }
74    if let Some(d) = &route.http.description {
75        op.insert("description".into(), json!(d));
76    }
77    if !route.http.tags.is_empty() {
78        op.insert("tags".into(), json!(&route.http.tags));
79    }
80
81    let parameters = build_parameters(route);
82    if !parameters.is_empty() {
83        op.insert("parameters".into(), Value::Array(parameters));
84    }
85
86    if let Some(request_body) = build_request_body(route) {
87        op.insert("requestBody".into(), request_body);
88    }
89
90    op.insert("responses".into(), build_responses(route));
91
92    if let Some(auth) = &route.http.auth
93        && !matches!(auth, AuthRequirement::None)
94        && let Some(name) = scheme_names.get(auth)
95    {
96        op.insert("security".into(), json!([{ name.as_str(): [] }]));
97    }
98
99    Value::Object(op)
100}
101
102fn build_parameters(route: &SqlRoute) -> Vec<Value> {
103    let mut out = Vec::new();
104    let parameter_schema = &route.metadata["parameter_schema"];
105    let properties = parameter_schema.get("properties").and_then(Value::as_object);
106    let Some(properties) = properties else {
107        return out;
108    };
109    let required: std::collections::HashSet<&str> = parameter_schema
110        .get("required")
111        .and_then(Value::as_array)
112        .map(|arr| arr.iter().filter_map(Value::as_str).collect())
113        .unwrap_or_default();
114
115    for (name, schema) in properties {
116        let location = match route.param_locations.get(name) {
117            Some(HttpParamBinding::Path) => "path",
118            Some(HttpParamBinding::Query) => "query",
119            Some(HttpParamBinding::Header) => "header",
120            _ => continue,
121        };
122        let is_required = location == "path" || required.contains(name.as_str());
123        let mut p = Map::new();
124        p.insert("name".into(), json!(name));
125        p.insert("in".into(), json!(location));
126        p.insert("required".into(), json!(is_required));
127        p.insert("schema".into(), schema.clone());
128        out.push(Value::Object(p));
129    }
130    out
131}
132
133fn build_request_body(route: &SqlRoute) -> Option<Value> {
134    let request_schema = route.metadata.get("request_schema")?;
135    if request_schema.is_null() {
136        return None;
137    }
138    Some(json!({
139        "required": true,
140        "content": {
141            "application/json": { "schema": request_schema }
142        }
143    }))
144}
145
146fn build_responses(route: &SqlRoute) -> Value {
147    let mut responses = Map::new();
148    let response_schema = route.metadata.get("response_schema").cloned().unwrap_or(Value::Null);
149    let codes: Vec<u16> = if route.http.status_codes.is_empty() {
150        vec![route.default_status]
151    } else {
152        route.http.status_codes.clone()
153    };
154    for (idx, code) in codes.iter().enumerate() {
155        let is_primary = idx == 0;
156        let mut body = Map::new();
157        body.insert("description".into(), json!(describe_status(*code)));
158        if is_primary && !response_schema.is_null() && *code != 204 {
159            body.insert(
160                "content".into(),
161                json!({ "application/json": { "schema": response_schema.clone() } }),
162            );
163        }
164        responses.insert(code.to_string(), Value::Object(body));
165    }
166    Value::Object(responses)
167}
168
169const fn describe_status(code: u16) -> &'static str {
170    match code {
171        200 => "OK",
172        201 => "Created",
173        202 => "Accepted",
174        204 => "No Content",
175        400 => "Bad Request",
176        401 => "Unauthorized",
177        403 => "Forbidden",
178        404 => "Not Found",
179        409 => "Conflict",
180        422 => "Unprocessable Entity",
181        500 => "Internal Server Error",
182        _ => "Response",
183    }
184}
185
186fn collect_security_schemes(
187    routes: &[SqlRoute],
188) -> (Map<String, Value>, std::collections::BTreeMap<AuthRequirement, String>) {
189    let mut schemes = Map::new();
190    let mut name_for = std::collections::BTreeMap::new();
191    for route in routes {
192        let Some(auth) = &route.http.auth else { continue };
193        if matches!(auth, AuthRequirement::None) {
194            continue;
195        }
196        if name_for.contains_key(auth) {
197            continue;
198        }
199        let name = match auth {
200            AuthRequirement::None => unreachable!(),
201            AuthRequirement::Bearer { format: None } => "bearerAuth".to_string(),
202            AuthRequirement::Bearer { format: Some(f) } => format!("bearer{}", f.to_uppercase()),
203            AuthRequirement::ApiKey { location, name } => {
204                format!("apiKey_{}_{}", location_short(*location), name.replace('-', "_"))
205            }
206        };
207        let scheme_value = match auth {
208            AuthRequirement::None => unreachable!(),
209            AuthRequirement::Bearer { format } => {
210                let mut s = Map::new();
211                s.insert("type".into(), json!("http"));
212                s.insert("scheme".into(), json!("bearer"));
213                if let Some(f) = format {
214                    s.insert("bearerFormat".into(), json!(f));
215                }
216                Value::Object(s)
217            }
218            AuthRequirement::ApiKey { location, name } => json!({
219                "type": "apiKey",
220                "in": location_str(*location),
221                "name": name,
222            }),
223        };
224        schemes.insert(name.clone(), scheme_value);
225        name_for.insert(auth.clone(), name);
226    }
227    (schemes, name_for)
228}
229
230const fn location_short(loc: ApiKeyLocation) -> &'static str {
231    match loc {
232        ApiKeyLocation::Header => "h",
233        ApiKeyLocation::Query => "q",
234        ApiKeyLocation::Cookie => "c",
235    }
236}
237
238const fn location_str(loc: ApiKeyLocation) -> &'static str {
239    match loc {
240        ApiKeyLocation::Header => "header",
241        ApiKeyLocation::Query => "query",
242        ApiKeyLocation::Cookie => "cookie",
243    }
244}
245
246const fn method_key(m: HttpMethod) -> &'static str {
247    match m {
248        HttpMethod::Get => "get",
249        HttpMethod::Post => "post",
250        HttpMethod::Put => "put",
251        HttpMethod::Patch => "patch",
252        HttpMethod::Delete => "delete",
253        HttpMethod::Head => "head",
254        HttpMethod::Options => "options",
255    }
256}
257
258#[cfg(test)]
259mod tests {
260    use super::*;
261    use crate::sql::neutral_to_json_schema::BuildOptions;
262    use crate::sql::route::route_from_query;
263    use scythe_core::analyzer::{AnalyzedColumn, AnalyzedParam, AnalyzedQuery};
264    use scythe_core::catalog::Catalog;
265    use scythe_core::parser::{CustomAnnotation, QueryCommand};
266
267    fn empty_catalog() -> Catalog {
268        Catalog::from_ddl(&[]).unwrap()
269    }
270
271    fn get_user_query() -> AnalyzedQuery {
272        AnalyzedQuery {
273            name: "GetUser".to_string(),
274            command: QueryCommand::One,
275            sql: "SELECT id, email FROM users WHERE id = $1".to_string(),
276            columns: vec![
277                AnalyzedColumn {
278                    name: "id".into(),
279                    neutral_type: "int64".into(),
280                    nullable: false,
281                },
282                AnalyzedColumn {
283                    name: "email".into(),
284                    neutral_type: "string".into(),
285                    nullable: false,
286                },
287            ],
288            params: vec![AnalyzedParam {
289                name: "id".into(),
290                neutral_type: "int64".into(),
291                nullable: false,
292                position: 1,
293            }],
294            deprecated: None,
295            source_table: Some("users".into()),
296            composites: vec![],
297            enums: vec![],
298            optional_params: vec![],
299            group_by: None,
300            custom: vec![
301                CustomAnnotation {
302                    name: "http".into(),
303                    value: "GET /users/{id}".into(),
304                    line: 1,
305                },
306                CustomAnnotation {
307                    name: "http_auth".into(),
308                    value: "bearer:jwt".into(),
309                    line: 2,
310                },
311                CustomAnnotation {
312                    name: "http_status".into(),
313                    value: "200,404".into(),
314                    line: 3,
315                },
316                CustomAnnotation {
317                    name: "http_tags".into(),
318                    value: "users".into(),
319                    line: 4,
320                },
321                CustomAnnotation {
322                    name: "http_summary".into(),
323                    value: "Fetch a user".into(),
324                    line: 5,
325                },
326            ],
327        }
328    }
329
330    fn create_user_query() -> AnalyzedQuery {
331        AnalyzedQuery {
332            name: "CreateUser".to_string(),
333            command: QueryCommand::ExecRows,
334            sql: "INSERT INTO users (email) VALUES ($1)".to_string(),
335            columns: vec![],
336            params: vec![AnalyzedParam {
337                name: "email".into(),
338                neutral_type: "string".into(),
339                nullable: false,
340                position: 1,
341            }],
342            deprecated: None,
343            source_table: None,
344            composites: vec![],
345            enums: vec![],
346            optional_params: vec![],
347            group_by: None,
348            custom: vec![
349                CustomAnnotation {
350                    name: "http".into(),
351                    value: "POST /users".into(),
352                    line: 1,
353                },
354                CustomAnnotation {
355                    name: "http_auth".into(),
356                    value: "bearer:jwt".into(),
357                    line: 2,
358                },
359                CustomAnnotation {
360                    name: "http_status".into(),
361                    value: "201".into(),
362                    line: 3,
363                },
364            ],
365        }
366    }
367
368    fn build_two_routes() -> Vec<SqlRoute> {
369        let opts = BuildOptions::default();
370        let r1 = route_from_query(&get_user_query(), &empty_catalog(), &opts)
371            .unwrap()
372            .unwrap();
373        let r2 = route_from_query(&create_user_query(), &empty_catalog(), &opts)
374            .unwrap()
375            .unwrap();
376        vec![r1, r2]
377    }
378
379    #[test]
380    fn emits_openapi_3_1_header() {
381        let routes = build_two_routes();
382        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("test", "0.1.0"));
383        assert_eq!(spec["openapi"], "3.1.0");
384        assert_eq!(spec["info"]["title"], "test");
385        assert_eq!(spec["info"]["version"], "0.1.0");
386    }
387
388    #[test]
389    fn groups_methods_under_shared_path() {
390        let routes = build_two_routes();
391        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
392        assert!(spec["paths"]["/users"]["post"].is_object());
393        assert!(spec["paths"]["/users/{id}"]["get"].is_object());
394    }
395
396    #[test]
397    fn operation_carries_operation_id_summary_tags() {
398        let routes = build_two_routes();
399        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
400        let op = &spec["paths"]["/users/{id}"]["get"];
401        assert_eq!(op["operationId"], "GetUser");
402        assert_eq!(op["summary"], "Fetch a user");
403        assert_eq!(op["tags"], json!(["users"]));
404    }
405
406    #[test]
407    fn path_parameter_emitted() {
408        let routes = build_two_routes();
409        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
410        let params = spec["paths"]["/users/{id}"]["get"]["parameters"].as_array().unwrap();
411        assert_eq!(params.len(), 1);
412        assert_eq!(params[0]["name"], "id");
413        assert_eq!(params[0]["in"], "path");
414        assert_eq!(params[0]["required"], true);
415    }
416
417    #[test]
418    fn post_carries_request_body() {
419        let routes = build_two_routes();
420        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
421        let body = &spec["paths"]["/users"]["post"]["requestBody"];
422        assert_eq!(body["required"], true);
423        assert!(body["content"]["application/json"]["schema"]["properties"]["email"].is_object());
424    }
425
426    #[test]
427    fn responses_keyed_by_status_codes() {
428        let routes = build_two_routes();
429        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
430        let resp = &spec["paths"]["/users/{id}"]["get"]["responses"];
431        assert!(resp["200"].is_object());
432        assert!(resp["404"].is_object());
433    }
434
435    #[test]
436    fn primary_response_includes_schema() {
437        let routes = build_two_routes();
438        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
439        let primary = &spec["paths"]["/users/{id}"]["get"]["responses"]["200"];
440        assert!(primary["content"]["application/json"]["schema"]["properties"]["id"].is_object());
441    }
442
443    #[test]
444    fn registers_bearer_security_scheme_once() {
445        let routes = build_two_routes();
446        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
447        let schemes = &spec["components"]["securitySchemes"];
448        assert_eq!(schemes.as_object().unwrap().len(), 1);
449        let (_name, scheme) = schemes.as_object().unwrap().iter().next().unwrap();
450        assert_eq!(scheme["type"], "http");
451        assert_eq!(scheme["scheme"], "bearer");
452        assert_eq!(scheme["bearerFormat"], "jwt");
453    }
454
455    #[test]
456    fn operations_reference_security_scheme() {
457        let routes = build_two_routes();
458        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
459        let op = &spec["paths"]["/users/{id}"]["get"];
460        let sec = op["security"].as_array().unwrap();
461        assert_eq!(sec.len(), 1);
462        let scheme_name = sec[0].as_object().unwrap().keys().next().unwrap();
463        assert!(spec["components"]["securitySchemes"][scheme_name].is_object());
464    }
465
466    #[test]
467    fn no_204_response_carries_body() {
468        let mut q = create_user_query();
469        q.command = QueryCommand::Exec;
470        q.custom.retain(|a| a.name != "http_status");
471        let route = route_from_query(&q, &empty_catalog(), &BuildOptions::default())
472            .unwrap()
473            .unwrap();
474        let spec = openapi_from_routes(&[route], &OpenApiInfo::new("t", "1"));
475        let resp = &spec["paths"]["/users"]["post"]["responses"]["204"];
476        assert!(resp["content"].is_null());
477    }
478}