Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
32 changes: 25 additions & 7 deletions grace/generators/model_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -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"
Expand All @@ -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
Expand Down Expand Up @@ -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

Expand Down
38 changes: 34 additions & 4 deletions tests/generators/test_model_generator.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,6 +2,7 @@

import pytest

from grace.exceptions import ValidationError
from grace.generator import Generator
from grace.generators.model_generator import ModelGenerator

Expand All @@ -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

Expand All @@ -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")
Loading