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