Skip to content

Commit 5a4b2f3

Browse files
feat: Scala case class support (Doobie/Quill/Slick) as 11th format (v1.6.0)
- NEW: ScalaParser — parses case class definitions with Option[T], default values - NEW: ScalaGenerator — produces case classes with Option[T] for nullable columns - Format identifier: 'scala' in CLI, MCP, and convert API - Sample fixture: fixtures/sample.scala - 17 new tests (parsing, generation, cross-format: sql/prisma/ef) - 270/270 tests passing - Updated CLI format choices, MCP format list - Bumped to v1.6.0
1 parent af389a1 commit 5a4b2f3

7 files changed

Lines changed: 552 additions & 2 deletions

File tree

‎fixtures/sample.scala‎

Lines changed: 19 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,19 @@
1+
import java.time.{Instant, LocalDate, LocalTime}
2+
import java.util.UUID
3+
4+
case class User(
5+
id: Int,
6+
name: String,
7+
email: String,
8+
role: String = "viewer",
9+
isActive: Boolean = true,
10+
createdAt: Instant
11+
)
12+
13+
case class Post(
14+
id: Int,
15+
title: String,
16+
content: Option[String],
17+
authorId: Int,
18+
status: String = "draft"
19+
)

‎src/schemaforge/cli.py‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -13,7 +13,7 @@
1313
from .mcp_server import mcp_command
1414

1515
# All supported format names (used for CLI choices and detection)
16-
_FORMATS = ["sql", "prisma", "drizzle", "typeorm", "django", "sqlalchemy", "alembic", "json_schema", "graphql", "ef"]
16+
_FORMATS = ["sql", "prisma", "drizzle", "typeorm", "django", "sqlalchemy", "alembic", "json_schema", "graphql", "ef", "scala"]
1717

1818

1919
@click.group()

‎src/schemaforge/convert.py‎

Lines changed: 3 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -24,6 +24,8 @@
2424
from .generators.graphql_generator import GraphQLGenerator
2525
from .parsers.ef_parser import EntityFrameworkParser
2626
from .generators.ef_generator import EntityFrameworkGenerator
27+
from .parsers.scala_parser import ScalaParser
28+
from .generators.scala_generator import ScalaGenerator
2729

2830
if TYPE_CHECKING:
2931
from .type_config import TypeConfig
@@ -40,6 +42,7 @@
4042
"json_schema": (JSONSchemaParser, JSONSchemaGenerator),
4143
"graphql": (GraphQLParser, GraphQLGenerator),
4244
"ef": (EntityFrameworkParser, EntityFrameworkGenerator),
45+
"scala": (ScalaParser, ScalaGenerator),
4346
}
4447

4548

Lines changed: 122 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,122 @@
1+
"""Generator: SchemaForge IR → Scala case class definitions.
2+
3+
Generates Scala case classes suitable for use with Doobie, Quill,
4+
or Slick. Handles Option[T] for nullable columns.
5+
"""
6+
from __future__ import annotations
7+
8+
from ..ir import Schema, Column, ColumnType
9+
from ._base import resolve_type
10+
11+
12+
class ScalaGenerator:
13+
"""Convert Schema IR to Scala case class definitions."""
14+
15+
_TYPE_MAP: dict[ColumnType, str] = {
16+
ColumnType.STRING: "String",
17+
ColumnType.INTEGER: "Int",
18+
ColumnType.FLOAT: "Double",
19+
ColumnType.BOOLEAN: "Boolean",
20+
ColumnType.DATETIME: "Instant",
21+
ColumnType.DATE: "LocalDate",
22+
ColumnType.TIME: "LocalTime",
23+
ColumnType.TEXT: "String",
24+
ColumnType.BLOB: "Array[Byte]",
25+
ColumnType.JSON: "String",
26+
ColumnType.UUID: "UUID",
27+
ColumnType.ENUM: "String",
28+
ColumnType.DECIMAL: "BigDecimal",
29+
}
30+
31+
def __init__(self, type_config=None):
32+
"""Initialize generator with optional type config."""
33+
self.type_config = type_config
34+
35+
def generate(self, schema: Schema) -> str:
36+
"""Generate Scala case class definitions from schema IR."""
37+
parts: list[str] = [
38+
"// Auto-generated by SchemaForge",
39+
"// https://github.com/Coding-Dev-Tools/schemaforge",
40+
"",
41+
"import java.time.{Instant, LocalDate, LocalTime}",
42+
"import java.util.UUID",
43+
"import scala.math.BigDecimal",
44+
"",
45+
]
46+
47+
# Generate imports based on types used
48+
needed_imports: set[str] = set()
49+
for table in schema.tables:
50+
for col in table.columns:
51+
sa_type = resolve_type(col, self._TYPE_MAP)
52+
if sa_type.startswith("java."):
53+
needed_imports.add(sa_type.rsplit(".", 1)[0])
54+
55+
if needed_imports:
56+
needed_imports.discard("java.time") # Already imported
57+
needed_imports.discard("java.util") # Already imported
58+
for imp in sorted(needed_imports):
59+
parts.append(f"import {imp}")
60+
61+
if needed_imports:
62+
parts.append("")
63+
64+
for table in schema.tables:
65+
parts.append(self._generate_case_class(table))
66+
67+
return "\n".join(parts)
68+
69+
def _generate_case_class(self, table) -> str:
70+
"""Generate a single case class."""
71+
class_name = self._to_pascal(table.name)
72+
lines: list[str] = []
73+
lines.append(f"case class {class_name}(")
74+
75+
field_lines: list[str] = []
76+
for col in table.columns:
77+
field_lines.append(f" {self._field_def(col)},")
78+
79+
if field_lines:
80+
lines.append("\n".join(field_lines))
81+
else:
82+
lines.append(" // No fields")
83+
84+
lines.append(")")
85+
86+
return "\n".join(lines)
87+
88+
def _field_def(self, col: Column) -> str:
89+
"""Generate a case class field definition."""
90+
sa_type = resolve_type(col, self._TYPE_MAP)
91+
92+
# Handle nullable → Option[T]
93+
if col.nullable:
94+
scala_type = f"Option[{sa_type}]"
95+
else:
96+
scala_type = sa_type
97+
98+
# Default value
99+
default_str = ""
100+
if col.default is not None:
101+
if isinstance(col.default, bool):
102+
default_str = f" = {str(col.default).lower()}"
103+
elif isinstance(col.default, str) and col.default.startswith("fn:"):
104+
fn_val = col.default[3:]
105+
fn_upper = fn_val.upper().rstrip("()")
106+
if fn_upper in ("CURRENT_TIMESTAMP", "NOW", "INSTANT_NOW"):
107+
default_str = " = java.time.Instant.now()"
108+
elif fn_upper == "CURRENT_DATE":
109+
default_str = " = java.time.LocalDate.now()"
110+
elif fn_upper == "UUID" or fn_upper == "RANDOM_UUID":
111+
default_str = " = java.util.UUID.randomUUID()"
112+
elif isinstance(col.default, str):
113+
default_str = f' = "{col.default}"'
114+
elif isinstance(col.default, (int, float)):
115+
default_str = f" = {col.default}"
116+
117+
return f"{col.name}: {scala_type}{default_str}"
118+
119+
def _to_pascal(self, name: str) -> str:
120+
"""Convert snake_case to PascalCase."""
121+
parts = name.split("_")
122+
return "".join(p.capitalize() for p in parts if p)

‎src/schemaforge/mcp_server.py‎

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -25,7 +25,7 @@
2525

2626

2727
# All supported formats
28-
_FORMATS = ["sql", "prisma", "drizzle", "typeorm", "django", "sqlalchemy", "alembic", "json_schema", "graphql", "ef"]
28+
_FORMATS = ["sql", "prisma", "drizzle", "typeorm", "django", "sqlalchemy", "alembic", "json_schema", "graphql", "ef", "scala"]
2929
_FORMAT_DESCRIPTIONS = {
3030
"sql": "SQL DDL (Data Definition Language)",
3131
"prisma": "Prisma schema",
@@ -37,6 +37,7 @@
3737
"json_schema": "JSON Schema (draft 2020-12)",
3838
"graphql": "GraphQL SDL (Schema Definition Language)",
3939
"ef": "Entity Framework Core entities (C#)",
40+
"scala": "Scala case classes (Doobie/Quill/Slick)",
4041
}
4142

4243

Lines changed: 191 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,191 @@
1+
"""Parser: Scala case class definitions → SchemaForge IR.
2+
3+
Parses Scala case class definitions suitable for Doobie, Quill,
4+
or Slick into the internal schema representation.
5+
"""
6+
from __future__ import annotations
7+
8+
import re
9+
from typing import Any
10+
11+
from ..ir import Schema, Table, Column, ColumnType
12+
13+
14+
# Regex: extract case class
15+
_CASE_CLASS_RE = re.compile(
16+
r'(?:@(?:Entity|Table|Mapped)\s*(?:\([^)]*\))?\s*)?'
17+
r'(?:case\s+)?class\s+(\w+)'
18+
r'(?:\s*\(((?:[^()]|\([^()]*\))*)\))',
19+
re.MULTILINE,
20+
)
21+
22+
# Regex: extract field from case class parameter list
23+
_FIELD_RE = re.compile(
24+
r'\s*(\w+)\s*:\s*([^=,]+)(?:\s*=\s*([^,]+))?\s*,?\s*'
25+
)
26+
27+
# Map Scala types to ColumnType
28+
_SCALA_TYPE_MAP: dict[str, ColumnType] = {
29+
"Int": ColumnType.INTEGER,
30+
"Long": ColumnType.INTEGER,
31+
"Short": ColumnType.INTEGER,
32+
"Byte": ColumnType.INTEGER,
33+
"BigInt": ColumnType.INTEGER,
34+
"String": ColumnType.STRING,
35+
"Boolean": ColumnType.BOOLEAN,
36+
"Double": ColumnType.FLOAT,
37+
"Float": ColumnType.FLOAT,
38+
"BigDecimal": ColumnType.DECIMAL,
39+
"java.math.BigDecimal": ColumnType.DECIMAL,
40+
"Instant": ColumnType.DATETIME,
41+
"java.time.Instant": ColumnType.DATETIME,
42+
"LocalDateTime": ColumnType.DATETIME,
43+
"java.time.LocalDateTime": ColumnType.DATETIME,
44+
"OffsetDateTime": ColumnType.DATETIME,
45+
"java.time.OffsetDateTime": ColumnType.DATETIME,
46+
"ZonedDateTime": ColumnType.DATETIME,
47+
"java.time.ZonedDateTime": ColumnType.DATETIME,
48+
"LocalDate": ColumnType.DATE,
49+
"java.time.LocalDate": ColumnType.DATE,
50+
"LocalTime": ColumnType.TIME,
51+
"java.time.LocalTime": ColumnType.TIME,
52+
"UUID": ColumnType.UUID,
53+
"java.util.UUID": ColumnType.UUID,
54+
"java.util.Date": ColumnType.DATETIME,
55+
"DateTime": ColumnType.DATETIME,
56+
"org.joda.time.DateTime": ColumnType.DATETIME,
57+
"org.joda.time.LocalDate": ColumnType.DATE,
58+
"org.joda.time.LocalTime": ColumnType.TIME,
59+
}
60+
61+
62+
def _clean_scala_type(raw: str) -> tuple[str, bool]:
63+
"""Clean a Scala type string.
64+
65+
Returns (cleaned_type, is_optional) where is_optional is True for Option[T].
66+
"""
67+
raw = raw.strip()
68+
is_optional = False
69+
70+
# Handle Option[T]
71+
if raw.startswith("Option["):
72+
raw = raw[7:-1].strip()
73+
is_optional = True
74+
75+
# Handle List[T], Seq[T], Vector[T]
76+
for prefix in ("List[", "Seq[", "Vector[", "Set[", "List["):
77+
if raw.startswith(prefix) and raw.endswith(">"):
78+
raw = raw[len(prefix) : -1].strip()
79+
break
80+
81+
return raw, is_optional
82+
83+
84+
def _parse_default(value_str: str) -> Any:
85+
"""Parse a Scala default value expression."""
86+
val = value_str.strip().rstrip(",")
87+
88+
if val == "true":
89+
return True
90+
if val == "false":
91+
return False
92+
if val == "null" or val == "None":
93+
return None
94+
95+
# Quoted strings
96+
if (val.startswith('"') and val.endswith('"')) or \
97+
(val.startswith('"""') and val.endswith('"""')):
98+
inner = val.strip('"')
99+
# Handle interpolation: s"..."
100+
idx = inner.find("$")
101+
if idx >= 0:
102+
inner = inner[:idx]
103+
return inner
104+
105+
# Integer
106+
try:
107+
return int(val)
108+
except ValueError:
109+
pass
110+
111+
# Float
112+
try:
113+
return float(val)
114+
except ValueError:
115+
pass
116+
117+
# Function defaults (e.g., java.time.Instant.now())
118+
if "(" in val and val.endswith(")"):
119+
fn_name = val.split("(")[0].split(".")[-1]
120+
return f"fn:{fn_name}()"
121+
122+
return None
123+
124+
125+
class ScalaParser:
126+
"""Parse Scala case class definitions into Schema IR."""
127+
128+
def parse(self, text: str) -> Schema:
129+
"""Parse Scala source text into a Schema IR."""
130+
schema = Schema()
131+
132+
# Remove comments
133+
text = re.sub(r"//.*$", "", text, flags=re.MULTILINE)
134+
text = re.sub(r"/\*.*?\*/", "", text, flags=re.DOTALL)
135+
136+
for m in _CASE_CLASS_RE.finditer(text):
137+
class_name = m.group(1)
138+
params_str = m.group(2)
139+
140+
if not class_name or not params_str:
141+
continue
142+
143+
table = Table(name=self._to_snake(class_name))
144+
fields = _FIELD_RE.findall(params_str)
145+
146+
for field_name, field_type, default_str in fields:
147+
col = self._field_to_column(field_name, field_type.strip(), default_str.strip())
148+
if col:
149+
table.columns.append(col)
150+
151+
if table.columns:
152+
schema.tables.append(table)
153+
154+
return schema
155+
156+
def _field_to_column(self, name: str, raw_type: str, default_str: str) -> Column | None:
157+
"""Convert a Scala field to a Column IR."""
158+
clean_type, is_optional = _clean_scala_type(raw_type)
159+
160+
col_type = _SCALA_TYPE_MAP.get(clean_type, ColumnType.CUSTOM)
161+
type_args: dict[str, Any] = {}
162+
default_val: Any = None
163+
custom_type = ""
164+
165+
# Parse default value
166+
if default_str:
167+
default_val = _parse_default(default_str)
168+
169+
# Handle String length from default values / heuristics
170+
if col_type == ColumnType.STRING and isinstance(default_val, str):
171+
pass # No length hint from Scala
172+
173+
if col_type == ColumnType.CUSTOM:
174+
custom_type = clean_type
175+
176+
return Column(
177+
name=name,
178+
type=col_type,
179+
type_args=type_args,
180+
nullable=is_optional,
181+
unique=False,
182+
primary_key=False,
183+
default=default_val,
184+
custom_type=custom_type,
185+
)
186+
187+
def _to_snake(self, name: str) -> str:
188+
"""Convert PascalCase to snake_case."""
189+
result = re.sub(r"([A-Z])", r"_\1", name).lower().lstrip("_")
190+
result = re.sub(r"([a-z0-9])([A-Z])", r"\1_\2", result).lower()
191+
return result.strip("_")

0 commit comments

Comments
 (0)