mirror of
https://github.com/langgenius/dify.git
synced 2026-09-08 11:04:27 +08:00
test: use sqlite3 session in test_passport (#38763)
This commit is contained in:
parent
b820ccf086
commit
14ac8dfbf6
@ -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
|
|
||||||
|
|||||||
Loading…
Reference in New Issue
Block a user