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