deerflow-code/offline-backend-20260512/backend/tests/test_gauss_database.py
2026-09-07 18:24:55 +08:00

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")