Skip to content
Closed
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
42 changes: 39 additions & 3 deletions backend/app/models/delivery.py
Original file line number Diff line number Diff line change
Expand Up @@ -17,6 +17,7 @@
String,
Text,
event,
inspect,
or_,
)
from sqlalchemy.engine import Connection
Expand Down Expand Up @@ -290,20 +291,51 @@ class DeliveryAsset(LoopNode):


def adapt_loop_node_values_for_dialect(
values: dict[str, object], dialect_name: str
values: dict[str, object],
dialect_name: str,
non_nullable_attributes: set[str] | frozenset[str] | None = None,
) -> dict[str, object]:
"""Convert explicit nulls to sentinels required by the production schema."""
if dialect_name != "mysql":
return values
adapted = values.copy()
for attribute, default in _MYSQL_NON_NULL_DEFAULTS.items():
attributes = (
_MYSQL_NON_NULL_DEFAULTS.keys()
if non_nullable_attributes is None
else non_nullable_attributes
)
for attribute in attributes:
default = _MYSQL_NON_NULL_DEFAULTS[attribute]
if attribute in adapted and adapted[attribute] is None:
adapted[attribute] = (
default.copy() if isinstance(default, dict) else default
)
return adapted


def loop_node_non_nullable_attributes(connection: Connection) -> frozenset[str]:
"""Return model attributes backed by NOT NULL columns in this MySQL schema."""
if connection.dialect.name != "mysql":
return frozenset()
cache_key = "loop_node_non_nullable_attributes"
cached = connection.info.get(cache_key)
if isinstance(cached, frozenset):
return cached
columns = {
column["name"]: column
for column in inspect(connection).get_columns("loop_items")
}
attributes = frozenset(
attribute
for attribute in _MYSQL_NON_NULL_DEFAULTS
if not columns[getattr(LoopNode, attribute).property.columns[0].name][
"nullable"
]
)
connection.info[cache_key] = attributes
return attributes


def loop_datetime_is_unset(column: object) -> object:
"""Match unset datetimes in both nullable and sentinel schemas."""
return or_(column.is_(None), column == _MYSQL_UNSET_DATETIME)
Expand All @@ -322,6 +354,10 @@ def _populate_mysql_non_null_defaults(
values = {
attribute: getattr(target, attribute) for attribute in _MYSQL_NON_NULL_DEFAULTS
}
adapted = adapt_loop_node_values_for_dialect(values, connection.dialect.name)
adapted = adapt_loop_node_values_for_dialect(
values,
connection.dialect.name,
loop_node_non_nullable_attributes(connection),
)
for attribute, value in adapted.items():
setattr(target, attribute, value)
6 changes: 5 additions & 1 deletion backend/app/services/loop_items/service.py
Original file line number Diff line number Diff line change
Expand Up @@ -28,6 +28,7 @@
adapt_loop_node_values_for_dialect,
loop_datetime_is_unset,
loop_datetime_value_is_unset,
loop_node_non_nullable_attributes,
)
from app.models.resource_member import MemberStatus, ResourceMember
from app.models.share_link import ResourceType
Expand Down Expand Up @@ -564,7 +565,9 @@ def update(
# its new lane instead of an arbitrary stale position.
updates["sort_order"] = 0
updates = adapt_loop_node_values_for_dialect(
updates, db.get_bind().dialect.name
updates,
db.get_bind().dialect.name,
loop_node_non_nullable_attributes(db.connection()),
)
updated = (
db.query(LoopItem)
Expand Down Expand Up @@ -834,6 +837,7 @@ def _advance_task_started_item(db: Session, item_id: str) -> None:
updates = adapt_loop_node_values_for_dialect(
{"status": "in_progress", "completed_at": None},
db.get_bind().dialect.name,
loop_node_non_nullable_attributes(db.connection()),
)
db.query(LoopItem).filter(
LoopItem.id == item_id,
Expand Down
16 changes: 16 additions & 0 deletions backend/tests/schemas/test_delivery.py
Original file line number Diff line number Diff line change
Expand Up @@ -54,3 +54,19 @@ def test_loop_item_update_adapts_nulls_only_for_mysql() -> None:
"completed_at": datetime(1970, 1, 1, 0, 0, 1),
}
assert sqlite_values == values


def test_loop_item_update_preserves_nullable_mysql_columns() -> None:
values = {"parent_id": None, "due_at": None, "completed_at": None}

adapted = adapt_loop_node_values_for_dialect(
values,
"mysql",
frozenset({"completed_at"}),
)

assert adapted == {
"parent_id": None,
"due_at": None,
"completed_at": datetime(1970, 1, 1, 0, 0, 1),
}
Loading