mvp
This commit is contained in:
@@ -0,0 +1,280 @@
|
||||
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
|
||||
|
||||
|
||||
@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",
|
||||
}
|
||||
|
||||
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"
|
||||
|
||||
fields.append(
|
||||
Field(
|
||||
name=field_name,
|
||||
type=field_type,
|
||||
description=field_schema.get("description"),
|
||||
imports=type_info.imports,
|
||||
)
|
||||
)
|
||||
|
||||
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
|
||||
Reference in New Issue
Block a user