|
17 | 17 | from sqlalchemy.exc import InvalidRequestError, SQLAlchemyError |
18 | 18 | from sqlalchemy.ext.asyncio import AsyncSession |
19 | 19 | from sqlalchemy.orm import InstrumentedAttribute, Mapped, Session, mapped_column |
| 20 | +from sqlalchemy.sql.selectable import ForUpdateArg |
20 | 21 | from sqlalchemy.types import TypeEngine |
21 | 22 |
|
22 | 23 | from advanced_alchemy import base |
@@ -378,6 +379,89 @@ async def test_sqlalchemy_repo_get_member( |
378 | 379 | mock_repo.session.commit.assert_not_called() # pyright: ignore[reportFunctionMemberAccess] |
379 | 380 |
|
380 | 381 |
|
| 382 | +async def test_sqlalchemy_repo_get_with_for_update( |
| 383 | + mock_repo: SQLAlchemyAsyncRepository[Any], |
| 384 | + mocker: MockerFixture, |
| 385 | +) -> None: |
| 386 | + """Ensure FOR UPDATE options are applied when requested.""" |
| 387 | + |
| 388 | + statement = MagicMock() |
| 389 | + statement.options.return_value = statement |
| 390 | + statement.execution_options.return_value = statement |
| 391 | + statement.with_for_update.return_value = statement |
| 392 | + mock_repo.statement = statement |
| 393 | + |
| 394 | + mocker.patch.object(mock_repo, "_get_loader_options", return_value=([], False)) |
| 395 | + mocker.patch.object(mock_repo, "_get_base_stmt", return_value=statement) |
| 396 | + mocker.patch.object(mock_repo, "_apply_filters", return_value=statement) |
| 397 | + mocker.patch.object(mock_repo, "_filter_select_by_kwargs", return_value=statement) |
| 398 | + execute_result = MagicMock() |
| 399 | + execute_result.scalar_one_or_none.return_value = MagicMock() |
| 400 | + execute = mocker.patch.object(mock_repo, "_execute", return_value=execute_result) |
| 401 | + |
| 402 | + instance = await maybe_async(mock_repo.get("instance-id", with_for_update=True)) |
| 403 | + |
| 404 | + assert instance is execute_result.scalar_one_or_none.return_value |
| 405 | + statement.with_for_update.assert_called_once_with() |
| 406 | + execute.assert_called_once_with(statement, uniquify=False) |
| 407 | + |
| 408 | + |
| 409 | +async def test_sqlalchemy_repo_get_with_for_update_dict( |
| 410 | + mock_repo: SQLAlchemyAsyncRepository[Any], |
| 411 | + mocker: MockerFixture, |
| 412 | +) -> None: |
| 413 | + statement = MagicMock() |
| 414 | + statement.options.return_value = statement |
| 415 | + statement.execution_options.return_value = statement |
| 416 | + statement.with_for_update.return_value = statement |
| 417 | + mock_repo.statement = statement |
| 418 | + |
| 419 | + mocker.patch.object(mock_repo, "_get_loader_options", return_value=([], False)) |
| 420 | + mocker.patch.object(mock_repo, "_get_base_stmt", return_value=statement) |
| 421 | + mocker.patch.object(mock_repo, "_apply_filters", return_value=statement) |
| 422 | + mocker.patch.object(mock_repo, "_filter_select_by_kwargs", return_value=statement) |
| 423 | + execute_result = MagicMock() |
| 424 | + execute_result.scalar_one_or_none.return_value = MagicMock() |
| 425 | + mocker.patch.object(mock_repo, "_execute", return_value=execute_result) |
| 426 | + |
| 427 | + await maybe_async( |
| 428 | + mock_repo.get( |
| 429 | + "instance-id", |
| 430 | + with_for_update={"nowait": True, "read": False}, |
| 431 | + ) |
| 432 | + ) |
| 433 | + |
| 434 | + statement.with_for_update.assert_called_once_with(nowait=True, read=False) |
| 435 | + |
| 436 | + |
| 437 | +async def test_sqlalchemy_repo_get_with_for_update_arg( |
| 438 | + mock_repo: SQLAlchemyAsyncRepository[Any], |
| 439 | + mocker: MockerFixture, |
| 440 | +) -> None: |
| 441 | + statement = MagicMock() |
| 442 | + statement.options.return_value = statement |
| 443 | + statement.execution_options.return_value = statement |
| 444 | + statement.with_for_update.return_value = statement |
| 445 | + mock_repo.statement = statement |
| 446 | + |
| 447 | + mocker.patch.object(mock_repo, "_get_loader_options", return_value=([], False)) |
| 448 | + mocker.patch.object(mock_repo, "_get_base_stmt", return_value=statement) |
| 449 | + mocker.patch.object(mock_repo, "_apply_filters", return_value=statement) |
| 450 | + mocker.patch.object(mock_repo, "_filter_select_by_kwargs", return_value=statement) |
| 451 | + execute_result = MagicMock() |
| 452 | + execute_result.scalar_one_or_none.return_value = MagicMock() |
| 453 | + mocker.patch.object(mock_repo, "_execute", return_value=execute_result) |
| 454 | + |
| 455 | + await maybe_async( |
| 456 | + mock_repo.get( |
| 457 | + "instance-id", |
| 458 | + with_for_update=ForUpdateArg(nowait=True, key_share=True), |
| 459 | + ) |
| 460 | + ) |
| 461 | + |
| 462 | + statement.with_for_update.assert_called_once_with(nowait=True, read=False, skip_locked=False, key_share=True) |
| 463 | + |
| 464 | + |
381 | 465 | async def test_sqlalchemy_repo_get_one_member( |
382 | 466 | mock_repo: SQLAlchemyAsyncRepository[Any], |
383 | 467 | monkeypatch: MonkeyPatch, |
|
0 commit comments