mirror of
https://github.com/langgenius/dify.git
synced 2026-08-15 04:59:46 +08:00
test: use SQLite sessions in controllers console app (#39096)
Co-authored-by: autofix-ci[bot] <114827586+autofix-ci[bot]@users.noreply.github.com>
This commit is contained in:
parent
9e95e1302e
commit
046df1260a
@ -1,6 +1,5 @@
|
||||
from __future__ import annotations
|
||||
|
||||
from contextlib import nullcontext
|
||||
from datetime import UTC, datetime
|
||||
from inspect import unwrap
|
||||
from types import SimpleNamespace
|
||||
@ -8,37 +7,45 @@ from types import SimpleNamespace
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from pydantic import ValidationError
|
||||
from sqlalchemy import Engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from controllers.console.app import conversation_variables as conversation_variables_module
|
||||
from factories import variable_factory
|
||||
from graphon.variables.types import SegmentType
|
||||
from models import ConversationVariable
|
||||
|
||||
|
||||
def test_get_conversation_variables_returns_paginated_response(app: Flask, monkeypatch: pytest.MonkeyPatch) -> None:
|
||||
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
|
||||
def test_get_conversation_variables_returns_paginated_response(
|
||||
app: Flask,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_engine: Engine,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
api = conversation_variables_module.ConversationVariablesApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
created_at = datetime(2026, 1, 1, tzinfo=UTC)
|
||||
updated_at = datetime(2026, 1, 2, tzinfo=UTC)
|
||||
row = SimpleNamespace(
|
||||
created_at=created_at,
|
||||
updated_at=updated_at,
|
||||
to_variable=lambda: SimpleNamespace(
|
||||
model_dump=lambda: {
|
||||
"id": "var-1",
|
||||
"name": "my_var",
|
||||
"value_type": "string",
|
||||
"value": "value",
|
||||
"description": "desc",
|
||||
}
|
||||
),
|
||||
)
|
||||
session = SimpleNamespace(scalars=lambda _stmt: SimpleNamespace(all=lambda: [row]))
|
||||
monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=object()))
|
||||
monkeypatch.setattr(
|
||||
conversation_variables_module,
|
||||
"sessionmaker",
|
||||
lambda *_args, **_kwargs: SimpleNamespace(begin=lambda: nullcontext(session)),
|
||||
variable = variable_factory.build_conversation_variable_from_mapping(
|
||||
{
|
||||
"id": "var-1",
|
||||
"name": "my_var",
|
||||
"value_type": SegmentType.STRING,
|
||||
"value": "value",
|
||||
"description": "desc",
|
||||
}
|
||||
)
|
||||
row = ConversationVariable.from_variable(app_id="app-1", conversation_id="conv-1", variable=variable)
|
||||
row.created_at = created_at
|
||||
row.updated_at = updated_at
|
||||
sqlite_session.add(row)
|
||||
sqlite_session.commit()
|
||||
sqlite_session.expire(row)
|
||||
expected_created_at = int(row.created_at.timestamp())
|
||||
expected_updated_at = int(row.updated_at.timestamp())
|
||||
monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=sqlite_engine))
|
||||
|
||||
with app.test_request_context(
|
||||
"/console/api/apps/app-1/conversation-variables",
|
||||
@ -52,36 +59,32 @@ def test_get_conversation_variables_returns_paginated_response(app: Flask, monke
|
||||
assert response["total"] == 1
|
||||
assert response["has_more"] is False
|
||||
assert response["data"][0]["id"] == "var-1"
|
||||
assert response["data"][0]["created_at"] == int(created_at.timestamp())
|
||||
assert response["data"][0]["updated_at"] == int(updated_at.timestamp())
|
||||
assert response["data"][0]["created_at"] == expected_created_at
|
||||
assert response["data"][0]["updated_at"] == expected_updated_at
|
||||
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(ConversationVariable,)], indirect=True)
|
||||
def test_get_conversation_variables_normalizes_value_type_and_value(
|
||||
app: Flask, monkeypatch: pytest.MonkeyPatch
|
||||
app: Flask,
|
||||
monkeypatch: pytest.MonkeyPatch,
|
||||
sqlite_engine: Engine,
|
||||
sqlite_session: Session,
|
||||
) -> None:
|
||||
api = conversation_variables_module.ConversationVariablesApi()
|
||||
method = unwrap(api.get)
|
||||
|
||||
row = SimpleNamespace(
|
||||
created_at=None,
|
||||
updated_at=None,
|
||||
to_variable=lambda: SimpleNamespace(
|
||||
model_dump=lambda: {
|
||||
"id": "var-2",
|
||||
"name": "my_var_2",
|
||||
"value_type": SegmentType.INTEGER,
|
||||
"value": 42,
|
||||
"description": None,
|
||||
}
|
||||
),
|
||||
)
|
||||
session = SimpleNamespace(scalars=lambda _stmt: SimpleNamespace(all=lambda: [row]))
|
||||
monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=object()))
|
||||
monkeypatch.setattr(
|
||||
conversation_variables_module,
|
||||
"sessionmaker",
|
||||
lambda *_args, **_kwargs: SimpleNamespace(begin=lambda: nullcontext(session)),
|
||||
variable = variable_factory.build_conversation_variable_from_mapping(
|
||||
{
|
||||
"id": "var-2",
|
||||
"name": "my_var_2",
|
||||
"value_type": SegmentType.INTEGER,
|
||||
"value": 42,
|
||||
"description": "",
|
||||
}
|
||||
)
|
||||
sqlite_session.add(ConversationVariable.from_variable(app_id="app-1", conversation_id="conv-1", variable=variable))
|
||||
sqlite_session.commit()
|
||||
monkeypatch.setattr(conversation_variables_module, "db", SimpleNamespace(engine=sqlite_engine))
|
||||
|
||||
with app.test_request_context(
|
||||
"/console/api/apps/app-1/conversation-variables",
|
||||
|
||||
@ -5,6 +5,9 @@ from unittest.mock import PropertyMock, patch
|
||||
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
from werkzeug.exceptions import NotFound
|
||||
|
||||
from controllers.console import console_ns
|
||||
from controllers.console.app.mcp_server import (
|
||||
@ -13,19 +16,42 @@ from controllers.console.app.mcp_server import (
|
||||
AppMCPServerResponse,
|
||||
)
|
||||
from controllers.console.wraps import RBACPermission, RBACResourceScope
|
||||
from models.enums import AppMCPServerStatus
|
||||
from models.model import AppMCPServer
|
||||
|
||||
|
||||
class _ValidatedResponse:
|
||||
def __init__(self, payload):
|
||||
def __init__(self, payload: dict[str, str]) -> None:
|
||||
self._payload = payload
|
||||
|
||||
def model_dump(self, mode="json"):
|
||||
def model_dump(self, mode: str = "json") -> dict[str, str]:
|
||||
return self._payload
|
||||
|
||||
|
||||
def _server(
|
||||
*,
|
||||
tenant_id: str = "tenant-1",
|
||||
app_id: str = "app-1",
|
||||
name: str = "Demo App",
|
||||
description: str = "Description",
|
||||
parameters: str = "{}",
|
||||
status: AppMCPServerStatus = AppMCPServerStatus.ACTIVE,
|
||||
server_code: str = "server-code",
|
||||
) -> AppMCPServer:
|
||||
return AppMCPServer(
|
||||
tenant_id=tenant_id,
|
||||
app_id=app_id,
|
||||
name=name,
|
||||
description=description,
|
||||
parameters=parameters,
|
||||
status=status,
|
||||
server_code=server_code,
|
||||
)
|
||||
|
||||
|
||||
class TestAppMCPServerResponse:
|
||||
def test_parameters_json_string_parsed(self):
|
||||
data = {
|
||||
def test_parameters_json_string_parsed(self) -> None:
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -36,8 +62,8 @@ class TestAppMCPServerResponse:
|
||||
resp = AppMCPServerResponse.model_validate(data)
|
||||
assert resp.parameters == {"key": "value"}
|
||||
|
||||
def test_parameters_invalid_json_returns_original(self):
|
||||
data = {
|
||||
def test_parameters_invalid_json_returns_original(self) -> None:
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -48,8 +74,8 @@ class TestAppMCPServerResponse:
|
||||
resp = AppMCPServerResponse.model_validate(data)
|
||||
assert resp.parameters == "not-valid-json"
|
||||
|
||||
def test_parameters_dict_passthrough(self):
|
||||
data = {
|
||||
def test_parameters_dict_passthrough(self) -> None:
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -60,8 +86,8 @@ class TestAppMCPServerResponse:
|
||||
resp = AppMCPServerResponse.model_validate(data)
|
||||
assert resp.parameters == {"already": "parsed"}
|
||||
|
||||
def test_parameters_json_array_parsed(self):
|
||||
data = {
|
||||
def test_parameters_json_array_parsed(self) -> None:
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -72,9 +98,9 @@ class TestAppMCPServerResponse:
|
||||
resp = AppMCPServerResponse.model_validate(data)
|
||||
assert resp.parameters == ["a", "b"]
|
||||
|
||||
def test_timestamps_normalized(self):
|
||||
def test_timestamps_normalized(self) -> None:
|
||||
dt = datetime.datetime(2024, 1, 1, 0, 0, 0, tzinfo=datetime.UTC)
|
||||
data = {
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -88,8 +114,8 @@ class TestAppMCPServerResponse:
|
||||
assert resp.created_at == int(dt.timestamp())
|
||||
assert resp.updated_at == int(dt.timestamp())
|
||||
|
||||
def test_timestamps_none(self):
|
||||
data = {
|
||||
def test_timestamps_none(self) -> None:
|
||||
data: dict[str, object] = {
|
||||
"id": "s1",
|
||||
"name": "test",
|
||||
"server_code": "code",
|
||||
@ -103,16 +129,18 @@ class TestAppMCPServerResponse:
|
||||
|
||||
|
||||
class TestAppMCPServerController:
|
||||
def test_get_returns_empty_dict_when_server_missing(self):
|
||||
@pytest.mark.parametrize("sqlite_session", [(AppMCPServer,)], indirect=True)
|
||||
def test_get_returns_empty_dict_when_server_missing(self, sqlite_session: Session) -> None:
|
||||
api = AppMCPServerController()
|
||||
method = unwrap(api.get)
|
||||
|
||||
with patch("controllers.console.app.mcp_server.db.session.scalar", return_value=None):
|
||||
with patch("controllers.console.app.mcp_server.db.session", sqlite_session):
|
||||
response = method(api, app_model=SimpleNamespace(id="app-1"))
|
||||
|
||||
assert response == {}
|
||||
|
||||
def test_post_returns_201(self):
|
||||
@pytest.mark.parametrize("sqlite_session", [(AppMCPServer,)], indirect=True)
|
||||
def test_post_returns_201(self, sqlite_session: Session) -> None:
|
||||
api = AppMCPServerController()
|
||||
method = unwrap(api.post)
|
||||
payload = {"parameters": {"timeout": 30}}
|
||||
@ -122,47 +150,35 @@ class TestAppMCPServerController:
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch("controllers.console.app.mcp_server.db.session.add"),
|
||||
patch("controllers.console.app.mcp_server.db.session.commit"),
|
||||
patch("controllers.console.app.mcp_server.db.session", sqlite_session),
|
||||
patch("controllers.console.app.mcp_server.AppMCPServer.generate_server_code", return_value="server-code"),
|
||||
patch(
|
||||
"controllers.console.app.mcp_server.AppMCPServerResponse.model_validate",
|
||||
return_value=_ValidatedResponse({"id": "server-1"}),
|
||||
),
|
||||
):
|
||||
response, status_code = method(
|
||||
api, "tenant-1", app_model=SimpleNamespace(id="app-1", name="Demo App", description="App description")
|
||||
)
|
||||
|
||||
assert response == {"id": "server-1"}
|
||||
server = sqlite_session.scalar(select(AppMCPServer))
|
||||
assert server is not None
|
||||
assert response["server_code"] == "server-code"
|
||||
assert response["parameters"] == {"timeout": 30}
|
||||
assert status_code == 201
|
||||
|
||||
def test_put_binds_server_lookup_to_app_ref(self):
|
||||
@pytest.mark.parametrize("sqlite_session", [(AppMCPServer,)], indirect=True)
|
||||
def test_put_updates_server_for_app(self, sqlite_session: Session) -> None:
|
||||
api = AppMCPServerController()
|
||||
method = unwrap(api.put)
|
||||
payload = {"id": "server-1", "description": "Updated", "parameters": {"timeout": 30}, "status": "active"}
|
||||
app = Flask(__name__)
|
||||
app.config["TESTING"] = True
|
||||
server = SimpleNamespace(
|
||||
id="server-1",
|
||||
tenant_id="tenant-1",
|
||||
app_id="app-1",
|
||||
name="Old",
|
||||
description="Old",
|
||||
parameters="{}",
|
||||
status="active",
|
||||
)
|
||||
server = _server(name="Old", description="Old")
|
||||
server.id = "server-1"
|
||||
sqlite_session.add(server)
|
||||
sqlite_session.commit()
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch("controllers.console.app.mcp_server.db.session.scalar", return_value=server) as scalar,
|
||||
patch("controllers.console.app.mcp_server.db.session.get") as get_mock,
|
||||
patch("controllers.console.app.mcp_server.db.session.commit") as commit,
|
||||
patch(
|
||||
"controllers.console.app.mcp_server.AppMCPServerResponse.model_validate",
|
||||
return_value=_ValidatedResponse({"id": "server-1"}),
|
||||
),
|
||||
patch("controllers.console.app.mcp_server.db.session", sqlite_session),
|
||||
):
|
||||
response = method(
|
||||
api,
|
||||
@ -171,18 +187,58 @@ class TestAppMCPServerController:
|
||||
),
|
||||
)
|
||||
|
||||
stmt = scalar.call_args.args[0]
|
||||
compiled = stmt.compile()
|
||||
statement = str(compiled)
|
||||
assert "app_mcp_servers.id" in statement
|
||||
assert "app_mcp_servers.tenant_id" in statement
|
||||
assert "app_mcp_servers.app_id" in statement
|
||||
assert payload["id"] in compiled.params.values()
|
||||
assert "tenant-1" in compiled.params.values()
|
||||
assert "app-1" in compiled.params.values()
|
||||
get_mock.assert_not_called()
|
||||
commit.assert_called_once()
|
||||
assert response == {"id": "server-1"}
|
||||
sqlite_session.expire_all()
|
||||
updated_server = sqlite_session.get(AppMCPServer, "server-1")
|
||||
assert updated_server is not None
|
||||
assert response["id"] == "server-1"
|
||||
assert updated_server.description == "Updated"
|
||||
|
||||
@pytest.mark.parametrize("sqlite_session", [(AppMCPServer,)], indirect=True)
|
||||
@pytest.mark.parametrize(
|
||||
("foreign_tenant_id", "foreign_app_id"),
|
||||
[
|
||||
("tenant-2", "app-1"),
|
||||
("tenant-1", "app-2"),
|
||||
],
|
||||
)
|
||||
def test_put_scopes_server_lookup_to_complete_app_ref(
|
||||
self,
|
||||
sqlite_session: Session,
|
||||
foreign_tenant_id: str,
|
||||
foreign_app_id: str,
|
||||
) -> None:
|
||||
api = AppMCPServerController()
|
||||
method = unwrap(api.put)
|
||||
payload = {"id": "server-1", "description": "Updated", "parameters": {"timeout": 30}, "status": "active"}
|
||||
app = Flask(__name__)
|
||||
app.config["TESTING"] = True
|
||||
foreign_server = _server(
|
||||
tenant_id=foreign_tenant_id,
|
||||
app_id=foreign_app_id,
|
||||
name="Other",
|
||||
server_code="other-code",
|
||||
)
|
||||
foreign_server.id = "server-1"
|
||||
sqlite_session.add(foreign_server)
|
||||
sqlite_session.commit()
|
||||
|
||||
with (
|
||||
app.test_request_context("/", json=payload),
|
||||
patch.object(type(console_ns), "payload", new_callable=PropertyMock, return_value=payload),
|
||||
patch("controllers.console.app.mcp_server.db.session", sqlite_session),
|
||||
pytest.raises(NotFound),
|
||||
):
|
||||
method(
|
||||
api,
|
||||
app_model=SimpleNamespace(
|
||||
id="app-1", tenant_id="tenant-1", name="Demo App", description="App description"
|
||||
),
|
||||
)
|
||||
|
||||
sqlite_session.expire_all()
|
||||
unchanged_server = sqlite_session.get(AppMCPServer, "server-1")
|
||||
assert unchanged_server is not None
|
||||
assert unchanged_server.description == "Description"
|
||||
|
||||
|
||||
class TestAppMCPServerRefreshController:
|
||||
|
||||
Loading…
Reference in New Issue
Block a user