diff --git a/grace/generators/model_generator.py b/grace/generators/model_generator.py index 39e7400..a39cb33 100644 --- a/grace/generators/model_generator.py +++ b/grace/generators/model_generator.py @@ -4,9 +4,21 @@ from click.core import Argument from jinja2_strcase.jinja2_strcase import to_snake +from grace.exceptions import ValidationError from grace.generator import Generator from grace.generators.migration_generator import generate_migration +# Field annotations must be plain Python types Pydantic can build a schema for, +# not SQLAlchemy column type classes (e.g. `str`, not `String`) - see +# https://docs.pydantic.dev/latest/concepts/types/ for what's supported. +COLUMN_TYPES: dict[str, str] = { + "String": "str", + "Text": "str", + "Integer": "int", + "Float": "float", + "Boolean": "bool", +} + class ModelGenerator(Generator): NAME: str = "model" @@ -24,9 +36,7 @@ def generate(self, name: str, params: tuple[str]): a SQLAlchemy-style definition. You can specify column names and types during generation using the format `column_name:Type`. - Supported types are any valid SQLAlchemy column types - (e.g., `String`, `Integer`, `Boolean`, etc.). - See https://docs.sqlalchemy.org/en/20/core/types.html + Supported types: String, Text, Integer, Float, Boolean. Example: ```bash @@ -80,12 +90,20 @@ def extract_columns(self, params: tuple[str]) -> tuple[list, list]: types = [] for param in params: - name, type = param.split(":") + name, type_ = param.split(":") + + if type_ not in COLUMN_TYPES: + raise ValidationError( + f"Unsupported column type '{type_}' for '{name}'. " + f"Supported types: {', '.join(sorted(COLUMN_TYPES))}." + ) + + python_type = COLUMN_TYPES[type_] - if type not in types: - types.append(type) + if python_type not in types: + types.append(python_type) - columns.append((name, type)) + columns.append((name, python_type)) return columns, types diff --git a/tests/generators/test_model_generator.py b/tests/generators/test_model_generator.py index fbd3588..4b67afc 100644 --- a/tests/generators/test_model_generator.py +++ b/tests/generators/test_model_generator.py @@ -2,6 +2,7 @@ import pytest +from grace.exceptions import ValidationError from grace.generator import Generator from grace.generators.model_generator import ModelGenerator @@ -13,6 +14,18 @@ def generator(): return generator +def test_extract_columns_with_supported_types__expect_python_annotations(generator): + columns, types = generator.extract_columns(("message:String", "age:Integer")) + + assert columns == [("message", "str"), ("age", "int")] + assert types == ["str", "int"] + + +def test_extract_columns_with_unsupported_type__expect_validation_error(generator): + with pytest.raises(ValidationError, match="Unsupported column type 'DateTime'"): + generator.extract_columns(("created_at:DateTime",)) + + def test_generate_without_app__expect_value_error(generator): generator.app = None @@ -32,15 +45,32 @@ def test_generate_without_database__expect_no_file_and_warning( assert "no database configured" in caplog.text.lower() -def test_generate_with_database__expect_model_file_and_migration_generated( - mocker, generator -): +def test_generate__expect_model_file_with_python_type_annotations(mocker, generator): mock_generate_file = mocker.patch.object(Generator, "generate_file") + mocker.patch("grace.generators.model_generator.generate_migration") + + generator.generate("Greeting", ("message:String",)) + + mock_generate_file.assert_called_once() + variables = mock_generate_file.call_args.kwargs["variables"] + assert list(variables["model_columns"]) == ["message: str"] + assert variables["model_column_types"] == ["str"] + + +def test_generate__expect_migration_generated(mocker, generator): + mocker.patch.object(Generator, "generate_file") mock_generate_migration = mocker.patch( "grace.generators.model_generator.generate_migration" ) generator.generate("Greeting", ("message:String",)) - mock_generate_file.assert_called_once() mock_generate_migration.assert_called_once_with(generator.app, "Create Greeting") + + +def test_validate_valid_name__expect_true(generator): + assert generator.validate("Greeting") + + +def test_validate_invalid_name__expect_false(generator): + assert not generator.validate("greeting")