test: use sqlite3 session in test_passport (#38763)

This commit is contained in:
Asuka Minato 2026-07-16 21:59:52 +09:00 committed by GitHub
parent b820ccf086
commit 14ac8dfbf6
No known key found for this signature in database
GPG Key ID: B5690EEEBB952194

View File

@ -1,8 +1,11 @@
from __future__ import annotations from __future__ import annotations
from types import SimpleNamespace from unittest.mock import patch
from uuid import NAMESPACE_URL, UUID, uuid5
import pytest import pytest
from sqlalchemy import select
from sqlalchemy.orm import Session
from werkzeug.exceptions import NotFound, Unauthorized from werkzeug.exceptions import NotFound, Unauthorized
from controllers.web.error import WebAppAuthRequiredError from controllers.web.error import WebAppAuthRequiredError
@ -12,9 +15,41 @@ from controllers.web.passport import (
exchange_token_for_existing_web_user, exchange_token_for_existing_web_user,
generate_session_id, generate_session_id,
) )
from models.enums import CustomizeTokenStrategy, EndUserType
from models.model import App, AppMode, EndUser, IconType, Site
from services.webapp_auth_service import WebAppAuthType from services.webapp_auth_service import WebAppAuthType
def _stable_uuid(value: str) -> str:
return str(uuid5(NAMESPACE_URL, value))
def _persist_webapp(session: Session, *, app_code: str = "code") -> tuple[App, Site]:
tenant_id = _stable_uuid(f"tenant:{app_code}")
app_model = App(
id=_stable_uuid(f"app:{app_code}"),
tenant_id=tenant_id,
name="Web App",
mode=AppMode.CHAT,
icon_type=IconType.EMOJI,
icon="chat",
icon_background="#FFFFFF",
enable_site=True,
enable_api=False,
)
site = Site(
id=_stable_uuid(f"site:{app_code}"),
app_id=app_model.id,
title="Web App Site",
default_language="en-US",
customize_token_strategy=CustomizeTokenStrategy.UUID,
code=app_code,
)
session.add_all([app_model, site])
session.commit()
return app_model, site
def test_decode_enterprise_webapp_user_id_none() -> None: def test_decode_enterprise_webapp_user_id_none() -> None:
assert decode_enterprise_webapp_user_id(None) is None assert decode_enterprise_webapp_user_id(None) is None
@ -31,70 +66,68 @@ def test_decode_enterprise_webapp_user_id_valid(monkeypatch: pytest.MonkeyPatch)
assert decode_enterprise_webapp_user_id("token") == decoded assert decode_enterprise_webapp_user_id("token") == decoded
def test_exchange_token_public_flow(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize("sqlite_session", [(App, Site)], indirect=True)
site = SimpleNamespace(id="s1", app_id="a1", code="code", status="normal") def test_exchange_token_public_flow(sqlite_session: Session) -> None:
app_model = SimpleNamespace(id="a1", status="normal", enable_site=True) app_model, site = _persist_webapp(sqlite_session)
call_state = {"calls": 0}
def _scalar_side_effect(*_args, **_kwargs):
call_state["calls"] += 1
return site if call_state["calls"] == 1 else app_model
db_session = SimpleNamespace(scalar=_scalar_side_effect)
monkeypatch.setattr("controllers.web.passport.db", SimpleNamespace(session=db_session))
monkeypatch.setattr("controllers.web.passport._exchange_for_public_app_token", lambda *_args, **_kwargs: "resp")
decoded = {"auth_type": "public"} decoded = {"auth_type": "public"}
result = exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.PUBLIC) with (
patch("controllers.web.passport.db.session", sqlite_session),
patch("controllers.web.passport._exchange_for_public_app_token", return_value="resp") as exchange_mock,
):
result = exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.PUBLIC)
assert result == "resp" assert result == "resp"
exchange_mock.assert_called_once_with(app_model, site, decoded)
def test_exchange_token_requires_external(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize("sqlite_session", [(App, Site)], indirect=True)
site = SimpleNamespace(id="s1", app_id="a1", code="code", status="normal") def test_exchange_token_requires_external(sqlite_session: Session) -> None:
app_model = SimpleNamespace(id="a1", status="normal", enable_site=True) _persist_webapp(sqlite_session)
call_state = {"calls": 0}
def _scalar_side_effect(*_args, **_kwargs):
call_state["calls"] += 1
return site if call_state["calls"] == 1 else app_model
db_session = SimpleNamespace(scalar=_scalar_side_effect)
monkeypatch.setattr("controllers.web.passport.db", SimpleNamespace(session=db_session))
decoded = {"auth_type": "internal"} decoded = {"auth_type": "internal"}
with pytest.raises(WebAppAuthRequiredError): with (
patch("controllers.web.passport.db.session", sqlite_session),
pytest.raises(WebAppAuthRequiredError),
):
exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.EXTERNAL) exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.EXTERNAL)
def test_exchange_token_missing_session_id(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize("sqlite_session", [(App, Site, EndUser)], indirect=True)
site = SimpleNamespace(id="s1", app_id="a1", code="code", status="normal") def test_exchange_token_missing_session_id(sqlite_session: Session) -> None:
app_model = SimpleNamespace(id="a1", status="normal", enable_site=True, tenant_id="t1") _persist_webapp(sqlite_session)
call_state = {"calls": 0}
def _scalar_side_effect(*_args, **_kwargs):
call_state["calls"] += 1
if call_state["calls"] == 1:
return site
if call_state["calls"] == 2:
return app_model
return None
db_session = SimpleNamespace(scalar=_scalar_side_effect, add=lambda *_a, **_k: None, commit=lambda: None)
monkeypatch.setattr("controllers.web.passport.db", SimpleNamespace(session=db_session))
decoded = {"auth_type": "internal"} decoded = {"auth_type": "internal"}
with pytest.raises(NotFound): with (
patch("controllers.web.passport.db.session", sqlite_session),
pytest.raises(NotFound),
):
exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.INTERNAL) exchange_token_for_existing_web_user("code", decoded, WebAppAuthType.INTERNAL)
assert sqlite_session.scalars(select(EndUser)).all() == []
def test_generate_session_id(monkeypatch: pytest.MonkeyPatch) -> None: @pytest.mark.parametrize("sqlite_session", [(EndUser,)], indirect=True)
counts = [1, 0] def test_generate_session_id(sqlite_session: Session) -> None:
collision_id = _stable_uuid("session:collision")
generated_id = _stable_uuid("session:generated")
sqlite_session.add(
EndUser(
id=_stable_uuid("end-user:collision"),
tenant_id=_stable_uuid("tenant:collision"),
type=EndUserType.BROWSER,
name="Existing User",
session_id=collision_id,
)
)
sqlite_session.commit()
def _scalar(*_args, **_kwargs): with (
return counts.pop(0) patch("controllers.web.passport.db.session", sqlite_session),
patch(
"controllers.web.passport.uuid.uuid4",
side_effect=[UUID(collision_id), UUID(generated_id)],
),
):
session_id = generate_session_id()
db_session = SimpleNamespace(scalar=_scalar) assert session_id == generated_id
monkeypatch.setattr("controllers.web.passport.db", SimpleNamespace(session=db_session))
session_id = generate_session_id()
assert session_id