| """ |
| Custom fixtures for atom_agent_endpoints integration tests. |
| |
| This conftest provides a simplified db_session fixture that avoids |
| the NoReferencedTableError issue when running multiple tests. |
| """ |
|
|
| import os |
| import sys |
| import tempfile |
| from pathlib import Path |
|
|
| |
| os.environ["TESTING"] = "1" |
|
|
| import pytest |
| from fastapi.testclient import TestClient |
| from sqlalchemy import create_engine, exc |
| from sqlalchemy.orm import Session, sessionmaker |
|
|
| |
| sys.path.insert(0, str(Path(__file__).parent.parent.parent)) |
|
|
| from main_api_app import app |
| from core.database import Base |
| from core.models import AgentRegistry, AgentExecution, AgentFeedback |
|
|
|
|
| @pytest.fixture(scope="function") |
| def db_session(): |
| """ |
| Create a fresh in-memory database for each test. |
| |
| Simplified version that avoids sorted_tables to prevent NoReferencedTableError. |
| """ |
| |
| fd, db_path = tempfile.mkstemp(suffix='.db') |
| os.close(fd) |
|
|
| engine = create_engine( |
| f"sqlite:///{db_path}", |
| connect_args={"check_same_thread": False}, |
| echo=False |
| ) |
|
|
| |
| engine._test_db_path = db_path |
|
|
| |
| |
| try: |
| Base.metadata.create_all(engine, checkfirst=True) |
| except exc.NoReferencedTableError: |
| |
| for table in Base.metadata.tables.values(): |
| try: |
| table.create(engine, checkfirst=True) |
| except exc.NoReferencedTableError: |
| |
| continue |
| except Exception as e: |
| |
| if "already exists" not in str(e).lower() and "duplicate" not in str(e).lower(): |
| raise |
|
|
| |
| TestingSessionLocal = sessionmaker(autocommit=False, autoflush=False, bind=engine) |
| session = TestingSessionLocal() |
|
|
| yield session |
|
|
| |
| session.close() |
| engine.dispose() |
| |
| if hasattr(engine, '_test_db_path'): |
| try: |
| os.unlink(engine._test_db_path) |
| except Exception: |
| pass |
|
|
|
|
| @pytest.fixture(scope="function") |
| def client(db_session: Session): |
| """ |
| Create TestClient with dependency override for test database. |
| """ |
| from core.database import get_db |
|
|
| def _get_db(): |
| try: |
| yield db_session |
| finally: |
| pass |
|
|
| app.dependency_overrides[get_db] = _get_db |
|
|
| |
| def _mock_get_current_user(): |
| from tests.factories.user_factory import AdminUserFactory |
| import uuid |
|
|
| unique_id = str(uuid.uuid4())[:8] |
| email = f"test_{unique_id}@integration.com" |
|
|
| try: |
| from core.models import User |
| user = db_session.query(User).filter(User.email == email).first() |
| if user: |
| return user |
| except Exception: |
| pass |
|
|
| user = AdminUserFactory(email=email, _session=db_session) |
| db_session.commit() |
| db_session.refresh(user) |
| return user |
|
|
| |
| try: |
| from core.auth import get_current_user |
| app.dependency_overrides[get_current_user] = _mock_get_current_user |
| except ImportError: |
| pass |
|
|
| |
| for middleware in app.user_middleware: |
| if hasattr(middleware, 'cls') and middleware.cls.__name__ == 'TrustedHostMiddleware': |
| middleware.kwargs['allowed_hosts'] = ['testserver', 'localhost', '127.0.0.1', '0.0.0.0', '*'] |
| break |
|
|
| |
| test_client = TestClient(app, base_url="http://testserver") |
|
|
| yield test_client |
|
|
| app.dependency_overrides.clear() |
|
|