285 lines
6.4 KiB
Python
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
|