from __future__ import annotations
import json
import re
import sys
from collections import defaultdict
HEADER = "// @generated by scripts/codegen.py — do not edit. Regenerate with scripts/codegen.sh.\n"
REF = "#/components/schemas/"
RUST_KEYWORDS = set(
"as break const continue crate else enum extern false fn for if impl in let loop match mod move mut "
"pub ref return self Self static struct super trait true type unsafe use where while async await dyn "
"abstract become box do final macro override priv typeof unsized virtual yield try gen".split()
)
MODULE_ALLOWS = (
"#![allow(clippy::doc_lazy_continuation, clippy::doc_markdown, clippy::struct_excessive_bools, "
"rustdoc::broken_intra_doc_links, rustdoc::bare_urls, rustdoc::invalid_html_tags)]\n"
)
def die(msg: str) -> None:
raise SystemExit(f"codegen: {msg}")
def snake(name: str) -> str:
s = re.sub(r"(.)([A-Z][a-z]+)", r"\1_\2", name)
s = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", s)
return re.sub(r"_+", "_", s).lower()
def camel(snake_name: str) -> str:
parts = snake_name.split("_")
return parts[0] + "".join(p[:1].upper() + p[1:] for p in parts[1:])
def field_name(prop: str) -> str:
s = snake(prop)
return s + "_" if s in RUST_KEYWORDS else s
def serde_rename(prop: str) -> str | None:
return None if camel(field_name(prop)) == prop else prop
def module_of(schema_name: str) -> tuple[str, str]:
if schema_name.startswith("Entities."):
return "entities", schema_name[len("Entities."):]
if schema_name.startswith("Models."):
return "models", schema_name[len("Models."):]
return "types", schema_name.rsplit(".", 1)[-1]
def rust_path(schema_name: str) -> str:
module, name = module_of(schema_name)
return f"crate::{module}::{name}"
def rust_type(prop: dict, schemas: dict, where: str = "?") -> str:
if "$ref" in prop:
target = prop["$ref"]
if not target.startswith(REF):
die(f"{where}: unsupported $ref {target!r}")
return rust_path(target[len(REF):])
t = prop.get("type")
if t == "string":
return "String"
if t == "integer":
return "i64" if prop.get("format") == "int64" else "i32"
if t == "number":
return "f64"
if t == "boolean":
return "bool"
if t == "array":
items = prop.get("items")
if not isinstance(items, dict):
die(f"{where}: array without items")
return f"Vec<{rust_type(items, schemas, where + '[]')}>"
if t == "object" and "properties" not in prop and isinstance(prop.get("additionalProperties"), dict):
return f"HashMap<String, {rust_type(prop['additionalProperties'], schemas, where + '{}')}>"
if t is None and not any(k in prop for k in ("oneOf", "anyOf", "allOf", "not", "properties", "items")):
return "serde_json::Value"
die(f"{where}: unsupported schema construct {json.dumps(prop)[:120]}")
return ""
def direct_struct_refs(schema: dict) -> list[str]:
out = []
for prop in schema.get("properties", {}).values():
if "$ref" in prop:
out.append(prop["$ref"][len(REF):])
return out
def cycle_edges(schemas: dict) -> set[tuple[str, str]]:
graph = {name: [t for t in direct_struct_refs(s) if t in schemas and "enum" not in schemas[t]]
for name, s in schemas.items() if "enum" not in s}
def reaches(start: str, goal: str) -> bool:
seen, stack = set(), [start]
while stack:
n = stack.pop()
if n == goal:
return True
if n in seen:
continue
seen.add(n)
stack.extend(graph.get(n, []))
return False
return {(src, dst) for src, dsts in graph.items() for dst in dsts if reaches(dst, src)}
def doc_lines(text: str | None, indent: str = "") -> str:
if not text:
return ""
lines = [ln.rstrip() for ln in text.replace("\r", "").strip().split("\n")]
return "".join(f"{indent}/// {ln}\n" if ln else f"{indent}///\n" for ln in lines)
def render_enum(name: str, schema_name: str, schema: dict) -> str:
values = schema["enum"]
names = schema.get("x-enumNames") or []
out = doc_lines(schema.get("description"))
out += f"/// OpenAPI schema: `{schema_name}`\n"
if len(values) != len(names):
out += (
f"///\n/// The spec lists {len(names)} names for {len(values)} values (.NET enum aliases), so the\n"
"/// name/value pairing is unknowable and this is kept as a transparent integer.\n"
"#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, serde::Serialize, serde::Deserialize)]\n"
"#[serde(transparent)]\n"
f"pub struct {name}(pub i32);\n\n"
)
return out
pairs = sorted(zip(values, names), key=lambda p: p[0])
lowest = pairs[0][0]
out += (
"#[derive(Debug, Clone, Copy, PartialEq, Eq, Hash, Default, "
"serde_repr::Serialize_repr, serde_repr::Deserialize_repr)]\n#[repr(i32)]\n"
f"pub enum {name} {{\n"
)
for value, variant in pairs:
if not re.match(r"^[A-Za-z_][A-Za-z0-9_]*$", variant):
die(f"{schema_name}: enum name {variant!r} is not a Rust identifier")
if value == lowest:
out += " #[default]\n"
out += f" {variant} = {value},\n"
out += "}\n\n"
return out
def render_struct(name: str, schema_name: str, schema: dict, schemas: dict, boxed: set) -> str:
required = set(schema.get("required", []))
out = doc_lines(schema.get("description"))
out += f"/// OpenAPI schema: `{schema_name}`\n"
out += "#[derive(Debug, Clone, PartialEq, Default, serde::Serialize, serde::Deserialize)]\n"
out += '#[serde(rename_all = "camelCase")]\n'
out += f"pub struct {name} {{\n"
props = schema.get("properties", {})
for prop in sorted(props, key=field_name):
p = props[prop]
where = f"{schema_name}.{prop}"
ty = rust_type(p, schemas, where)
if "$ref" in p and (schema_name, p["$ref"][len(REF):]) in boxed:
ty = f"Box<{ty}>"
optional = p.get("nullable", False) or prop not in required
desc = p.get("description") or ""
if p.get("format") and p.get("type") == "string":
desc = (desc + "\n\n" if desc else "") + f"Format: `{p['format']}`."
if p.get("deprecated"):
desc = (desc + "\n\n" if desc else "") + "Deprecated by the API."
out += doc_lines(desc, " ")
attrs = []
rename = serde_rename(prop)
if rename is not None:
attrs.append(f'rename = "{rename}"')
if optional:
attrs.append('default, skip_serializing_if = "Option::is_none"')
ty = f"Option<{ty}>"
if attrs:
out += f" #[serde({', '.join(attrs)})]\n"
out += f" pub {field_name(prop)}: {ty},\n"
out += "}\n\n"
return out
def plan_modules(schemas: dict) -> dict[str, list[tuple[str, str]]]:
modules: dict[str, dict[str, str]] = defaultdict(dict)
for schema_name in schemas:
module, name = module_of(schema_name)
if name in modules[module]:
die(f"name collision in module `{module}`: {modules[module][name]} and {schema_name} both map to {name}")
modules[module][name] = schema_name
return {m: sorted(d.items()) for m, d in modules.items()}
def render_module(module: str, members: list[tuple[str, str]], schemas: dict, boxed: set) -> str:
out = HEADER + MODULE_ALLOWS + "#[allow(unused_imports)]\nuse std::collections::HashMap;\n\n"
for name, schema_name in members:
schema = schemas[schema_name]
if "enum" in schema:
if schema.get("type") != "integer":
die(f"{schema_name}: only integer enums are supported")
out += render_enum(name, schema_name, schema)
elif schema.get("type") == "object" or "properties" in schema:
out += render_struct(name, schema_name, schema, schemas, boxed)
else:
die(f"{schema_name}: top-level schema is neither object nor integer enum")
return out
def main(argv: list[str]) -> int:
if len(argv) != 3:
print(__doc__, file=sys.stderr)
return 2
spec_path, out_dir = argv[1], argv[2]
with open(spec_path, encoding="utf-8") as fh:
schemas = json.load(fh)["components"]["schemas"]
modules = plan_modules(schemas)
boxed = cycle_edges(schemas)
for module in ("entities", "models", "types"):
text = render_module(module, modules.get(module, []), schemas, boxed)
with open(f"{out_dir}/{module}.rs", "w", encoding="utf-8", newline="\n") as fh:
fh.write(text)
print(f"codegen: {module}.rs - {len(modules.get(module, []))} types")
with open(f"{out_dir}/mod.rs", "w", encoding="utf-8", newline="\n") as fh:
fh.write(HEADER + "pub mod entities;\npub mod models;\npub mod types;\n")
print(f"codegen: {len(boxed)} cycle edge(s) boxed")
return 0
if __name__ == "__main__":
sys.exit(main(sys.argv))