1use 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
34pub 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}