82 lines
2.8 KiB
Python
82 lines
2.8 KiB
Python
"""Regression coverage for the openGauss/GaussDB persistence backend."""
|
|
|
|
from __future__ import annotations
|
|
|
|
from datetime import UTC, datetime
|
|
|
|
import pytest
|
|
from pydantic import ValidationError
|
|
from sqlalchemy.engine import make_url
|
|
from sqlalchemy.schema import CreateIndex, CreateTable
|
|
|
|
from deerflow.config.database_config import DatabaseConfig
|
|
from deerflow.persistence.types import BeijingDateTime
|
|
|
|
|
|
@pytest.mark.parametrize(
|
|
("configured", "expected"),
|
|
[
|
|
(
|
|
"opengauss://user:pass@127.0.0.1:26000/deerflow",
|
|
"opengauss+asyncpg://user:pass@127.0.0.1:26000/deerflow",
|
|
),
|
|
(
|
|
"postgresql://user:pass@127.0.0.1:26000/deerflow?ssl=require",
|
|
"opengauss+asyncpg://user:pass@127.0.0.1:26000/deerflow?ssl=require",
|
|
),
|
|
(
|
|
"opengauss+asyncpg://user:pass@127.0.0.1:26000/deerflow",
|
|
"opengauss+asyncpg://user:pass@127.0.0.1:26000/deerflow",
|
|
),
|
|
],
|
|
)
|
|
def test_gauss_urls_always_select_the_opengauss_async_dialect(
|
|
configured: str,
|
|
expected: str,
|
|
) -> None:
|
|
config = DatabaseConfig(backend="gauss", gauss_url=configured)
|
|
assert config.app_sqlalchemy_url == expected
|
|
|
|
|
|
def test_gauss_backend_requires_a_supported_nonempty_url() -> None:
|
|
with pytest.raises(ValueError, match="gauss_url is required"):
|
|
_ = DatabaseConfig(backend="gauss").app_sqlalchemy_url
|
|
|
|
with pytest.raises(ValueError, match="must use opengauss"):
|
|
_ = DatabaseConfig(
|
|
backend="gauss",
|
|
gauss_url="mysql://user:pass@127.0.0.1/deerflow",
|
|
).app_sqlalchemy_url
|
|
|
|
with pytest.raises(ValidationError):
|
|
DatabaseConfig(backend="unknown")
|
|
|
|
|
|
def test_gauss_timestamps_keep_timezone_information() -> None:
|
|
dialect = type("GaussDialect", (), {"name": "opengauss"})()
|
|
value = datetime(2026, 9, 2, 8, 30, tzinfo=UTC)
|
|
column_type = BeijingDateTime()
|
|
|
|
assert column_type.process_bind_param(value, dialect) is value
|
|
assert column_type.process_result_value(value, dialect) is value
|
|
|
|
|
|
def test_official_async_dialect_compiles_the_complete_orm_schema() -> None:
|
|
"""Every registered table/index must compile before startup touches GaussDB."""
|
|
pytest.importorskip("opengauss_sqlalchemy")
|
|
|
|
import deerflow.persistence.models # noqa: F401 -- register ORM metadata
|
|
from deerflow.persistence.base import Base
|
|
|
|
dialect_cls = make_url("opengauss+asyncpg://user:pass@localhost/db").get_dialect(
|
|
_is_async=True
|
|
)
|
|
dialect = dialect_cls()
|
|
|
|
assert dialect.__class__.__module__.startswith("opengauss_sqlalchemy")
|
|
assert Base.metadata.sorted_tables
|
|
for table in Base.metadata.sorted_tables:
|
|
assert str(CreateTable(table).compile(dialect=dialect)).startswith("\nCREATE TABLE")
|
|
for index in table.indexes:
|
|
assert str(CreateIndex(index).compile(dialect=dialect)).startswith("CREATE")
|