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::build(|q| {
273            q.name = "GetUser".to_string();
274            q.command = QueryCommand::One;
275            q.sql = "SELECT id, email FROM users WHERE id = $1".to_string();
276            q.columns = vec![
277                AnalyzedColumn {
278                    name: "id".into(),
279                    neutral_type: "int64".into(),
280                    nullable: false,
281                    ..Default::default()
282                },
283                AnalyzedColumn {
284                    name: "email".into(),
285                    neutral_type: "string".into(),
286                    nullable: false,
287                    ..Default::default()
288                },
289            ];
290            q.params = vec![AnalyzedParam {
291                name: "id".into(),
292                neutral_type: "int64".into(),
293                nullable: false,
294                position: 1,
295                ..Default::default()
296            }];
297            q.deprecated = None;
298            q.source_table = Some("users".into());
299            q.composites = vec![];
300            q.enums = vec![];
301            q.optional_params = vec![];
302            q.group_by = None;
303            q.custom = vec![
304                CustomAnnotation {
305                    name: "http".into(),
306                    value: "GET /users/{id}".into(),
307                    line: 1,
308                    suggested_keyword: None,
309                },
310                CustomAnnotation {
311                    name: "http_auth".into(),
312                    value: "bearer:jwt".into(),
313                    line: 2,
314                    suggested_keyword: None,
315                },
316                CustomAnnotation {
317                    name: "http_status".into(),
318                    value: "200,404".into(),
319                    line: 3,
320                    suggested_keyword: None,
321                },
322                CustomAnnotation {
323                    name: "http_tags".into(),
324                    value: "users".into(),
325                    line: 4,
326                    suggested_keyword: None,
327                },
328                CustomAnnotation {
329                    name: "http_summary".into(),
330                    value: "Fetch a user".into(),
331                    line: 5,
332                    suggested_keyword: None,
333                },
334            ];
335        })
336    }
337
338    fn create_user_query() -> AnalyzedQuery {
339        AnalyzedQuery::build(|q| {
340            q.name = "CreateUser".to_string();
341            q.command = QueryCommand::ExecRows;
342            q.sql = "INSERT INTO users (email) VALUES ($1)".to_string();
343            q.columns = vec![];
344            q.params = vec![AnalyzedParam {
345                name: "email".into(),
346                neutral_type: "string".into(),
347                nullable: false,
348                position: 1,
349                ..Default::default()
350            }];
351            q.deprecated = None;
352            q.source_table = None;
353            q.composites = vec![];
354            q.enums = vec![];
355            q.optional_params = vec![];
356            q.group_by = None;
357            q.custom = vec![
358                CustomAnnotation {
359                    name: "http".into(),
360                    value: "POST /users".into(),
361                    line: 1,
362                    suggested_keyword: None,
363                },
364                CustomAnnotation {
365                    name: "http_auth".into(),
366                    value: "bearer:jwt".into(),
367                    line: 2,
368                    suggested_keyword: None,
369                },
370                CustomAnnotation {
371                    name: "http_status".into(),
372                    value: "201".into(),
373                    line: 3,
374                    suggested_keyword: None,
375                },
376            ];
377        })
378    }
379
380    fn build_two_routes() -> Vec<SqlRoute> {
381        let opts = BuildOptions::default();
382        let r1 = route_from_query(&get_user_query(), &empty_catalog(), &opts)
383            .unwrap()
384            .unwrap();
385        let r2 = route_from_query(&create_user_query(), &empty_catalog(), &opts)
386            .unwrap()
387            .unwrap();
388        vec![r1, r2]
389    }
390
391    #[test]
392    fn emits_openapi_3_1_header() {
393        let routes = build_two_routes();
394        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("test", "0.1.0"));
395        assert_eq!(spec["openapi"], "3.1.0");
396        assert_eq!(spec["info"]["title"], "test");
397        assert_eq!(spec["info"]["version"], "0.1.0");
398    }
399
400    #[test]
401    fn groups_methods_under_shared_path() {
402        let routes = build_two_routes();
403        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
404        assert!(spec["paths"]["/users"]["post"].is_object());
405        assert!(spec["paths"]["/users/{id}"]["get"].is_object());
406    }
407
408    #[test]
409    fn operation_carries_operation_id_summary_tags() {
410        let routes = build_two_routes();
411        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
412        let op = &spec["paths"]["/users/{id}"]["get"];
413        assert_eq!(op["operationId"], "GetUser");
414        assert_eq!(op["summary"], "Fetch a user");
415        assert_eq!(op["tags"], json!(["users"]));
416    }
417
418    #[test]
419    fn path_parameter_emitted() {
420        let routes = build_two_routes();
421        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
422        let params = spec["paths"]["/users/{id}"]["get"]["parameters"].as_array().unwrap();
423        assert_eq!(params.len(), 1);
424        assert_eq!(params[0]["name"], "id");
425        assert_eq!(params[0]["in"], "path");
426        assert_eq!(params[0]["required"], true);
427    }
428
429    #[test]
430    fn post_carries_request_body() {
431        let routes = build_two_routes();
432        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
433        let body = &spec["paths"]["/users"]["post"]["requestBody"];
434        assert_eq!(body["required"], true);
435        assert!(body["content"]["application/json"]["schema"]["properties"]["email"].is_object());
436    }
437
438    #[test]
439    fn responses_keyed_by_status_codes() {
440        let routes = build_two_routes();
441        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
442        let resp = &spec["paths"]["/users/{id}"]["get"]["responses"];
443        assert!(resp["200"].is_object());
444        assert!(resp["404"].is_object());
445    }
446
447    #[test]
448    fn primary_response_includes_schema() {
449        let routes = build_two_routes();
450        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
451        let primary = &spec["paths"]["/users/{id}"]["get"]["responses"]["200"];
452        assert!(primary["content"]["application/json"]["schema"]["properties"]["id"].is_object());
453    }
454
455    #[test]
456    fn registers_bearer_security_scheme_once() {
457        let routes = build_two_routes();
458        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
459        let schemes = &spec["components"]["securitySchemes"];
460        assert_eq!(schemes.as_object().unwrap().len(), 1);
461        let (_name, scheme) = schemes.as_object().unwrap().iter().next().unwrap();
462        assert_eq!(scheme["type"], "http");
463        assert_eq!(scheme["scheme"], "bearer");
464        assert_eq!(scheme["bearerFormat"], "jwt");
465    }
466
467    #[test]
468    fn operations_reference_security_scheme() {
469        let routes = build_two_routes();
470        let spec = openapi_from_routes(&routes, &OpenApiInfo::new("t", "1"));
471        let op = &spec["paths"]["/users/{id}"]["get"];
472        let sec = op["security"].as_array().unwrap();
473        assert_eq!(sec.len(), 1);
474        let scheme_name = sec[0].as_object().unwrap().keys().next().unwrap();
475        assert!(spec["components"]["securitySchemes"][scheme_name].is_object());
476    }
477
478    #[test]
479    fn no_204_response_carries_body() {
480        let mut q = create_user_query();
481        q.command = QueryCommand::Exec;
482        q.custom.retain(|a| a.name != "http_status");
483        let route = route_from_query(&q, &empty_catalog(), &BuildOptions::default())
484            .unwrap()
485            .unwrap();
486        let spec = openapi_from_routes(&[route], &OpenApiInfo::new("t", "1"));
487        let resp = &spec["paths"]["/users"]["post"]["responses"]["204"];
488        assert!(resp["content"].is_null());
489    }
490}