Files
sync-protocol/src/generator/languages/python.py
T

285 lines
6.4 KiB
Python

from dataclasses import dataclass
from pathlib import Path
from .base import BaseGenerator
@dataclass
class TypeInfo:
type: str
imports: set[str]
@dataclass
class Field:
name: str
type: str
description: str | None = None
imports: set[str] | None = None
optional: bool = False
@dataclass
class Model:
name: str
description: str
fields: list[Field]
class PythonGenerator(BaseGenerator):
TYPE_MAP = {
"string": "str",
"integer": "int",
"number": "float",
"boolean": "bool",
"object": "dict",
"float": "float",
}
FORMAT_MAP = {
"int64": "int",
}
def generate(self, output_dir: str | Path, make_all: bool, schema_file: str | None):
output_dir = Path(output_dir)
output_dir.mkdir(
parents=True,
exist_ok=True,
)
generated = set()
if make_all:
for schema_file in self.schemas:
self._generate_schema(
schema_file=schema_file,
output_dir=output_dir,
generated=generated,
)
return
self._generate_schema(
schema_file=schema_file,
output_dir=output_dir,
generated=generated,
)
def _generate_schema(
self,
schema_file: str,
output_dir: Path,
generated: set[str],
):
# Already generated
if schema_file in generated:
return
schema = self.schemas.get(schema_file)
if schema is None:
raise ValueError(f"Schema not found: {schema_file}")
# Mark as generated before resolving dependencies.
#
# This also prevents infinite recursion if schemas
# eventually reference each other.
generated.add(schema_file)
# First generate all referenced schemas.
for ref in self._find_refs(schema):
ref_file = self._ref_to_filename(ref)
self._generate_schema(
schema_file=ref_file,
output_dir=output_dir,
generated=generated,
)
# Now generate this schema itself.
model = self._parse_model(schema)
imports = self._collect_imports([model])
template = self.env.get_template("python.jinja2")
content = template.render(
model=model,
imports=sorted(imports),
)
module_name = self._schema_to_module(schema_file)
output_file = output_dir / f"{module_name}.py"
output_file.write_text(
content,
encoding="utf-8",
)
def _find_refs(self, schema: dict) -> list[str]:
"""
Recursively find all $ref values inside a schema.
"""
refs = []
if "$ref" in schema:
refs.append(schema["$ref"])
for value in schema.values():
if isinstance(value, dict):
refs.extend(self._find_refs(value))
elif isinstance(value, list):
for item in value:
if isinstance(item, dict):
refs.extend(self._find_refs(item))
return refs
def _ref_to_filename(self, ref: str) -> str:
"""
Convert:
operation.schema.json
into:
operation.schema.json
Also handles:
./operation.schema.json
"""
return Path(ref).name
def _parse_model(
self,
schema: dict,
) -> Model:
name = schema["title"]
fields = []
required = set(schema.get("required", []))
for (
field_name,
field_schema,
) in schema.get(
"properties",
{},
).items():
type_info = self._resolve_type(field_schema)
field_type = type_info.type
if field_name not in required:
field_type = f"{field_type} | None = None"
fields.append(
Field(
name=field_name,
type=field_type,
description=field_schema.get("description"),
imports=type_info.imports,
optional=field_name not in required,
)
)
fields.sort(key=lambda field: field.optional)
return Model(
name=name,
description=schema.get(
"description",
"",
),
fields=fields,
)
def _resolve_type(
self,
schema: dict,
) -> TypeInfo:
# $ref
if "$ref" in schema:
return self._resolve_ref(schema["$ref"])
# array
if schema.get("type") == "array":
item = self._resolve_type(schema["items"])
return TypeInfo(
type=f"list[{item.type}]",
imports=item.imports,
)
# format
schema_format = schema.get("format")
if schema_format in self.FORMAT_MAP:
return TypeInfo(
type=self.FORMAT_MAP[schema_format],
imports=set(),
)
# primitive
json_type = schema.get("type")
if json_type in self.TYPE_MAP:
return TypeInfo(
type=self.TYPE_MAP[json_type],
imports=set(),
)
raise ValueError(f"Unsupported schema: {schema}")
def _resolve_ref(
self,
ref: str,
) -> TypeInfo:
filename = self._ref_to_filename(ref)
schema = self.schemas.get(filename)
if schema is None:
raise ValueError(f"Referenced schema not found: {filename}")
model_name = schema["title"]
module_name = self._schema_to_module(filename)
return TypeInfo(
type=model_name,
imports={f"from .{module_name} import {model_name}"},
)
def _schema_to_module(
self,
filename: str,
) -> str:
return filename.removesuffix(".schema.json").replace("-", "_")
def _collect_imports(
self,
models: list[Model],
) -> set[str]:
imports = set()
for model in models:
for field in model.fields:
if field.imports:
imports.update(field.imports)
return imports