use std::collections::BTreeSet;
use heck::{ToShoutySnakeCase, ToSnakeCase, ToUpperCamelCase};
use sekkei::{OpenApiSpec, Operation, Parameter, Schema, ref_name};
#[must_use]
pub fn emit(spec: &OpenApiSpec, package: &str) -> String {
let mut imports: BTreeSet<String> = BTreeSet::new();
let mut body = String::new();
if let Some(c) = &spec.components {
for (name, schema) in &c.schemas {
emit_named(name, schema, &mut body, &mut imports);
}
}
emit_service(spec, package, &mut body, &mut imports);
let mut out = String::from("syntax = \"proto3\";\n\n");
out.push_str(&format!("package {package};\n\n"));
for imp in &imports {
out.push_str(&format!("import \"{imp}\";\n"));
}
if !imports.is_empty() {
out.push('\n');
}
out.push_str(&body);
out
}
const STRUCT: &str = "google.protobuf.Struct";
const EMPTY: &str = "google.protobuf.Empty";
fn field_type(schema: &Schema, imports: &mut BTreeSet<String>) -> (String, String) {
if schema.is_ref() {
return (
String::new(),
ref_name(schema.ref_path.as_deref().unwrap_or("")).to_upper_camel_case(),
);
}
if schema.is_array() {
let inner = schema.items.as_deref().cloned().unwrap_or_default();
let (_, ity) = field_type(&inner, imports);
return (String::from("repeated "), ity);
}
let ty = match schema.schema_type.as_deref() {
Some("integer") => if schema.format.as_deref() == Some("int32") {
"int32"
} else {
"int64"
}
.to_string(),
Some("number") => if schema.format.as_deref() == Some("float") {
"float"
} else {
"double"
}
.to_string(),
Some("string") => match schema.format.as_deref() {
Some("byte" | "binary") => "bytes".to_string(),
_ => "string".to_string(),
},
Some("boolean") => "bool".to_string(),
Some("object") => {
if let Some(ap) = &schema.additional_properties {
let (_, vty) = field_type(ap, imports);
format!("map<string, {vty}>")
} else {
imports.insert("google/protobuf/struct.proto".into());
STRUCT.to_string()
}
}
_ => {
imports.insert("google/protobuf/struct.proto".into());
STRUCT.to_string()
}
};
(String::new(), ty)
}
fn emit_named(name: &str, schema: &Schema, out: &mut String, imports: &mut BTreeSet<String>) {
let msg = name.to_upper_camel_case();
if schema.is_enum() {
emit_enum(&msg, schema, out);
} else if schema.is_object() || !schema.properties.is_empty() {
emit_message(&msg, schema, out, imports);
} else if schema.is_array() {
let (label, ity) = field_type(schema, imports);
out.push_str(&format!(
"message {msg} {{\n {label}{ity} items = 1;\n}}\n\n"
));
} else if schema.is_primitive() {
let (_, ity) = field_type(schema, imports);
out.push_str(&format!("message {msg} {{\n {ity} value = 1;\n}}\n\n"));
} else {
imports.insert("google/protobuf/struct.proto".into());
out.push_str(&format!("message {msg} {{\n {STRUCT} value = 1;\n}}\n\n"));
}
}
fn emit_enum(name: &str, schema: &Schema, out: &mut String) {
out.push_str(&format!("enum {name} {{\n"));
let prefix = name.to_shouty_snake_case();
out.push_str(&format!(" {prefix}_UNSPECIFIED = 0;\n"));
if let Some(values) = &schema.enum_values {
for (i, v) in values.iter().enumerate() {
if let Some(s) = v.as_str() {
out.push_str(&format!(
" {prefix}_{} = {};\n",
s.to_shouty_snake_case(),
i + 1
));
}
}
}
out.push_str("}\n\n");
}
fn emit_message(name: &str, schema: &Schema, out: &mut String, imports: &mut BTreeSet<String>) {
out.push_str(&format!("message {name} {{\n"));
for (i, (prop, pschema)) in schema.properties.iter().enumerate() {
let (label, ty) = field_type(pschema, imports);
let presence = if pschema.nullable && label.is_empty() {
"optional "
} else {
""
};
out.push_str(&format!(
" {presence}{label}{ty} {} = {};\n",
prop.to_snake_case(),
i + 1
));
}
out.push_str("}\n\n");
}
#[must_use]
pub fn service_name(package: &str) -> String {
package
.split('.')
.next()
.unwrap_or("Api")
.to_upper_camel_case()
}
pub struct RpcSig {
pub rpc: String,
pub method: String,
pub req_type: String,
pub resp_type: String,
}
#[must_use]
pub fn rpc_signatures(spec: &OpenApiSpec) -> Vec<RpcSig> {
let mut sink = String::new();
let mut imps = BTreeSet::new();
spec.all_operations()
.filter_map(|(_m, _p, op)| {
let op_id = op.operation_id.as_ref()?;
let rpc = op_id.to_upper_camel_case();
let resp_type = response_type(&rpc, op.success_response_schema(), &mut sink, &mut imps);
Some(RpcSig {
method: rpc.to_snake_case(),
req_type: format!("{rpc}Request"),
resp_type,
rpc,
})
})
.collect()
}
fn merged_params<'a>(spec: &'a OpenApiSpec, path: &str, op: &'a Operation) -> Vec<&'a Parameter> {
let mut out: Vec<&Parameter> = Vec::new();
if let Some(item) = spec.paths.get(path) {
for p in &item.parameters {
out.push(p);
}
}
for p in &op.parameters {
if let Some(slot) = out
.iter_mut()
.find(|x| x.name == p.name && x.location == p.location)
{
*slot = p;
} else {
out.push(p);
}
}
out
}
fn emit_service(
spec: &OpenApiSpec,
package: &str,
out: &mut String,
imports: &mut BTreeSet<String>,
) {
let service = service_name(package);
let mut rpcs = String::new();
let mut messages = String::new();
for (_method, path, op) in spec.all_operations() {
let Some(op_id) = &op.operation_id else {
continue;
};
let rpc = op_id.to_upper_camel_case();
let req_name = format!("{rpc}Request");
let mut field_no = 0usize;
let mut req_fields = String::new();
for p in merged_params(spec, &path, op) {
if matches!(p.location.as_str(), "path" | "query") {
let sch = p.schema.clone().unwrap_or_default();
let (label, ty) = field_type(&sch, imports);
field_no += 1;
req_fields.push_str(&format!(
" {label}{ty} {} = {};\n",
p.name.to_snake_case(),
field_no
));
}
}
if let Some(body) = op.json_body_schema() {
if body.is_ref() {
field_no += 1;
let ty = ref_name(body.ref_path.as_deref().unwrap_or("")).to_upper_camel_case();
req_fields.push_str(&format!(" {ty} body = {field_no};\n"));
} else if !body.properties.is_empty() {
for (prop, pschema) in &body.properties {
let (label, ty) = field_type(pschema, imports);
field_no += 1;
req_fields.push_str(&format!(
" {label}{ty} {} = {};\n",
prop.to_snake_case(),
field_no
));
}
} else {
imports.insert("google/protobuf/struct.proto".into());
field_no += 1;
req_fields.push_str(&format!(" {STRUCT} body = {field_no};\n"));
}
}
messages.push_str(&format!("message {req_name} {{\n{req_fields}}}\n\n"));
let resp_ty = response_type(&rpc, op.success_response_schema(), &mut messages, imports);
rpcs.push_str(&format!(" rpc {rpc}({req_name}) returns ({resp_ty});\n"));
}
out.push_str(&messages);
out.push_str(&format!("service {service} {{\n{rpcs}}}\n"));
}
fn response_type(
rpc: &str,
schema: Option<&Schema>,
messages: &mut String,
imports: &mut BTreeSet<String>,
) -> String {
let Some(schema) = schema else {
imports.insert("google/protobuf/empty.proto".into());
return EMPTY.to_string();
};
if schema.is_ref() {
return ref_name(schema.ref_path.as_deref().unwrap_or("")).to_upper_camel_case();
}
let resp = format!("{rpc}Response");
if schema.is_array() {
let (label, ity) = field_type(schema, imports);
messages.push_str(&format!(
"message {resp} {{\n {label}{ity} items = 1;\n}}\n\n"
));
} else if !schema.properties.is_empty() {
emit_message(&resp, schema, messages, imports);
} else {
imports.insert("google/protobuf/struct.proto".into());
messages.push_str(&format!("message {resp} {{\n {STRUCT} value = 1;\n}}\n\n"));
}
resp
}
#[cfg(test)]
mod tests {
use super::*;
const SPEC: &str = r##"
openapi: 3.0.3
info: { title: breathe control API, version: 0.1.0 }
paths:
/api/v1/catalog:
get:
operationId: catalogList
responses:
"200": { content: { application/json: { schema: { $ref: "#/components/schemas/Catalog" } } } }
/api/v1/bands/{kind}:
get:
operationId: bandList
parameters:
- { name: kind, in: path, required: true, schema: { type: string } }
- { name: namespace, in: query, schema: { type: string } }
responses:
"200": { content: { application/json: { schema: { type: array, items: { $ref: "#/components/schemas/Band" } } } } }
/api/v1/bands/{kind}/{namespace}/{name}/dry-run:
patch:
operationId: bandSetDryRun
parameters:
- { name: kind, in: path, schema: { type: string } }
- { name: namespace, in: path, schema: { type: string } }
- { name: name, in: path, schema: { type: string } }
requestBody:
content: { application/json: { schema: { type: object, properties: { dryRun: { type: boolean } } } } }
responses:
"200": { content: { application/json: { schema: { $ref: "#/components/schemas/Band" } } } }
components:
schemas:
BandKind: { type: string, enum: [memory, cpu, storage, arc, cgroup] }
BandStatus:
type: object
properties:
phase: { type: string }
lastChangeEpoch: { type: integer }
Band:
type: object
properties:
spec: { type: object }
status: { $ref: "#/components/schemas/BandStatus" }
Catalog:
type: object
properties:
dimensions: { type: array, items: { type: object } }
"##;
fn proto() -> String {
let spec: OpenApiSpec = serde_yaml_ng::from_str(SPEC).unwrap();
emit(&spec, "breathe.v1")
}
#[test]
fn header_package_and_syntax() {
let p = proto();
assert!(p.starts_with("syntax = \"proto3\";"));
assert!(p.contains("package breathe.v1;"));
}
#[test]
fn enum_maps_with_unspecified_zero() {
let p = proto();
assert!(p.contains("enum BandKind {"));
assert!(p.contains("BAND_KIND_UNSPECIFIED = 0;"));
assert!(p.contains("BAND_KIND_MEMORY = 1;"));
assert!(p.contains("BAND_KIND_CGROUP = 5;"));
}
#[test]
fn object_schema_becomes_typed_message_with_ref_and_scalar() {
let p = proto();
assert!(p.contains("message BandStatus {"));
assert!(p.contains("int64 last_change_epoch = 1;"));
assert!(p.contains("string phase = 2;"));
assert!(p.contains("BandStatus status ="));
}
#[test]
fn service_rpcs_with_synthesized_request_and_typed_response() {
let p = proto();
assert!(p.contains("service Breathe {"));
assert!(p.contains("rpc BandList(BandListRequest) returns (BandListResponse);"));
assert!(p.contains("repeated Band items = 1;"));
assert!(p.contains("rpc CatalogList(CatalogListRequest) returns (Catalog);"));
assert!(p.contains("rpc BandSetDryRun(BandSetDryRunRequest) returns (Band);"));
assert!(p.contains("bool dry_run ="));
assert!(p.contains("string kind ="));
}
#[test]
fn struct_import_only_when_used() {
let p = proto();
assert!(p.contains("import \"google/protobuf/struct.proto\";"));
assert!(p.contains("google.protobuf.Struct spec ="));
}
const PATH_LEVEL_SPEC: &str = r##"
openapi: 3.0.3
info: { title: t, version: 0.1.0 }
paths:
/api/v1/bands/{kind}/{namespace}/{name}:
parameters:
- { name: kind, in: path, required: true, schema: { $ref: "#/components/schemas/BandKind" } }
- { name: namespace, in: path, required: true, schema: { type: string } }
- { name: name, in: path, required: true, schema: { type: string } }
get:
operationId: bandGet
responses:
"200": { content: { application/json: { schema: { $ref: "#/components/schemas/Band" } } } }
patch:
operationId: bandPatch
requestBody:
content: { application/json: { schema: { $ref: "#/components/schemas/BandSpec" } } }
responses:
"200": { content: { application/json: { schema: { $ref: "#/components/schemas/Band" } } } }
components:
schemas:
BandKind: { type: string, enum: [memory, arc] }
BandSpec: { type: object, properties: { setpoint: { type: number } } }
Band: { type: object, properties: { kind: { type: string } } }
"##;
#[test]
fn nullable_field_becomes_proto3_optional() {
let spec_src = r##"
openapi: 3.0.3
info: { title: t, version: 0.1.0 }
paths: {}
components:
schemas:
DimensionSpec:
type: object
properties:
id: { type: string }
upstreamMirror: { type: string, nullable: true }
tags: { type: array, items: { type: string }, nullable: true }
"##;
let spec: OpenApiSpec = serde_yaml_ng::from_str(spec_src).unwrap();
let p = emit(&spec, "x.v1");
assert!(p.contains("optional string upstream_mirror ="));
assert!(p.contains("string id ="));
assert!(!p.contains("optional string id"));
assert!(p.contains("repeated string tags ="));
assert!(!p.contains("optional repeated"));
}
#[test]
fn path_level_params_merge_into_request_messages() {
let spec: OpenApiSpec = serde_yaml_ng::from_str(PATH_LEVEL_SPEC).unwrap();
let p = emit(&spec, "breathe.v1");
assert!(p.contains("message BandGetRequest {"));
assert!(p.contains("BandKind kind = 1;"));
assert!(p.contains("string namespace = 2;"));
assert!(p.contains("string name = 3;"));
assert!(p.contains("message BandPatchRequest {"));
assert!(p.contains("BandSpec body = 4;"));
assert!(p.contains("rpc BandPatch(BandPatchRequest) returns (Band);"));
}
}