diff --git a/api/controllers/inner_api/__init__.py b/api/controllers/inner_api/__init__.py index f47861cf274..986ebd29738 100644 --- a/api/controllers/inner_api/__init__.py +++ b/api/controllers/inner_api/__init__.py @@ -23,6 +23,7 @@ from .knowledge import retrieval as _knowledge_retrieval from .plugin import agent_config as _agent_config from .plugin import agent_drive as _agent_drive from .plugin import plugin as _plugin +from .workspace import plugin_model_providers as _plugin_model_providers from .workspace import workspace as _workspace api.add_namespace(inner_api_ns) @@ -35,6 +36,7 @@ __all__ = [ "_knowledge_retrieval", "_mail", "_plugin", + "_plugin_model_providers", "_runtime_credentials", "_workspace", "api", diff --git a/api/controllers/inner_api/workspace/plugin_model_providers.py b/api/controllers/inner_api/workspace/plugin_model_providers.py new file mode 100644 index 00000000000..50008a5bd82 --- /dev/null +++ b/api/controllers/inner_api/workspace/plugin_model_providers.py @@ -0,0 +1,39 @@ +from flask_restx import Resource +from pydantic import BaseModel, ConfigDict, Field + +from controllers.common.schema import register_schema_model +from controllers.console.wraps import setup_required +from controllers.inner_api import inner_api_ns +from controllers.inner_api.wraps import enterprise_inner_api_only +from core.plugin.plugin_service import PluginService + + +class InvalidatePluginModelProvidersCachePayload(BaseModel): + model_config = ConfigDict(extra="forbid") + + tenant_ids: list[str] = Field(default_factory=list, description="Workspace ids whose cache should be invalidated") + + +register_schema_model(inner_api_ns, InvalidatePluginModelProvidersCachePayload) + + +@inner_api_ns.route("/enterprise/workspace/plugin-model-providers/invalidate") +class EnterprisePluginModelProvidersCacheInvalidate(Resource): + @setup_required + @enterprise_inner_api_only + @inner_api_ns.doc( + "enterprise_invalidate_plugin_model_providers_cache", + responses={ + 200: "Cache invalidated", + 400: "Invalid request", + 401: "Unauthorized - invalid API key", + }, + ) + @inner_api_ns.expect(inner_api_ns.models[InvalidatePluginModelProvidersCachePayload.__name__]) + def post(self): + args = InvalidatePluginModelProvidersCachePayload.model_validate(inner_api_ns.payload or {}) + + for tenant_id in args.tenant_ids: + PluginService.invalidate_plugin_model_providers_cache(tenant_id) + + return {"result": "success"}, 200 diff --git a/api/tests/unit_tests/controllers/inner_api/workspace/test_plugin_model_providers.py b/api/tests/unit_tests/controllers/inner_api/workspace/test_plugin_model_providers.py new file mode 100644 index 00000000000..25902117ce5 --- /dev/null +++ b/api/tests/unit_tests/controllers/inner_api/workspace/test_plugin_model_providers.py @@ -0,0 +1,64 @@ +import inspect +from unittest.mock import call, patch + +import pytest +from flask import Flask +from pydantic import ValidationError + +from controllers.inner_api.workspace.plugin_model_providers import ( + EnterprisePluginModelProvidersCacheInvalidate, + InvalidatePluginModelProvidersCachePayload, +) + + +class TestInvalidatePluginModelProvidersCachePayload: + def test_valid_payload(self): + payload = InvalidatePluginModelProvidersCachePayload.model_validate( + {"tenant_ids": ["tenant-alpha", "tenant-beta"]} + ) + assert payload.tenant_ids == ["tenant-alpha", "tenant-beta"] + + def test_missing_tenant_ids_defaults_to_empty(self): + payload = InvalidatePluginModelProvidersCachePayload.model_validate({}) + assert payload.tenant_ids == [] + + def test_unknown_field_rejected(self): + with pytest.raises(ValidationError): + InvalidatePluginModelProvidersCachePayload.model_validate({"tenant_ids": ["tenant-alpha"], "generation": 7}) + + +class TestEnterprisePluginModelProvidersCacheInvalidate: + @pytest.fixture + def api_instance(self): + return EnterprisePluginModelProvidersCacheInvalidate() + + def _post(self, api_instance, app: Flask, payload): + unwrapped_post = inspect.unwrap(api_instance.post) + with app.test_request_context(): + with patch("controllers.inner_api.workspace.plugin_model_providers.inner_api_ns") as mock_ns: + mock_ns.payload = payload + return unwrapped_post(api_instance) + + @patch("controllers.inner_api.workspace.plugin_model_providers.PluginService") + def test_post_invalidates_once_per_tenant(self, mock_plugin_service, api_instance, app: Flask): + result = self._post(api_instance, app, {"tenant_ids": ["tenant-alpha", "tenant-beta"]}) + + assert result == ({"result": "success"}, 200) + assert mock_plugin_service.invalidate_plugin_model_providers_cache.call_args_list == [ + call("tenant-alpha"), + call("tenant-beta"), + ] + + @patch("controllers.inner_api.workspace.plugin_model_providers.PluginService") + def test_post_with_empty_list_is_a_no_op(self, mock_plugin_service, api_instance, app: Flask): + result = self._post(api_instance, app, {"tenant_ids": []}) + + assert result == ({"result": "success"}, 200) + mock_plugin_service.invalidate_plugin_model_providers_cache.assert_not_called() + + @patch("controllers.inner_api.workspace.plugin_model_providers.PluginService") + def test_post_with_missing_payload_is_a_no_op(self, mock_plugin_service, api_instance, app: Flask): + result = self._post(api_instance, app, None) + + assert result == ({"result": "success"}, 200) + mock_plugin_service.invalidate_plugin_model_providers_cache.assert_not_called()