Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
175 changes: 175 additions & 0 deletions aqueduct/gateway/tests/test_plugins_api.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,175 @@
from typing import ClassVar

from django.test import override_settings

from gateway.tests.test_endpoints import ChatCompletionsBase
from management.models import Snippet, SnippetType
from management.plugins import _plugin_class
from mock_api.mock_configs import MockConfig

AFTER_CALLS: list[bool] = []
EVENTS: list[str] = []
BODY_EXTRA_KEY = "custom_plugin_field"


BODY_MUTATING_PLUGIN = """\
from gateway.tests.test_plugins_api import BODY_EXTRA_KEY
class Mutator(PluginSnippet):
def before_request(self, request, token, body):
body[BODY_EXTRA_KEY] = "plugin-added"
return body
"""


BLOCKING_PLUGIN = """\
class Guard(PluginSnippet):
def before_request(self, request, token, body):
raise BlockedByPlugin("forbidden by guard", status=403)
"""


AFTER_PLUGIN = """\
from gateway.tests.test_plugins_api import AFTER_CALLS
class Audit(PluginSnippet):
def after_response(self, request, token, response):
AFTER_CALLS.append(True)
"""


PLUGIN_RECORD_A = """\
from gateway.tests.test_plugins_api import EVENTS
class A(PluginSnippet):
def before_request(self, request, token, body):
EVENTS.append("A")
return body
"""


PLUGIN_RECORD_B = """\
from gateway.tests.test_plugins_api import EVENTS
class B(PluginSnippet):
def before_request(self, request, token, body):
EVENTS.append("B")
return body
"""


PLUGIN_SET_STEP = """\
from gateway.tests.test_plugins_api import EVENTS
class SetStep(PluginSnippet):
def before_request(self, request, token, body):
body["step"] = "from-pipe"
return body
"""


PLUGIN_READ_STEP = """\
from gateway.tests.test_plugins_api import EVENTS
class ReadStep(PluginSnippet):
def before_request(self, request, token, body):
EVENTS.append(body.get("step"))
return body
"""


PLUGIN_BOOM = """\
class Boom(PluginSnippet):
def before_request(self, request, token, body):
raise RuntimeError("plugin exploded")
"""


def seed_plugin(name: str, code: str, order: int = 0) -> Snippet:
return Snippet.objects.create(
name=name, type=SnippetType.PLUGIN, active=True, order=order, code=code
)


@override_settings(TIKA_SERVER_URL=None)
class PluginGatewayIntegrationTest(ChatCompletionsBase):
MESSAGES: ClassVar[list[dict[str, str]]] = [{"role": "user", "content": "hello"}]

def setUp(self):
super().setUp()
_plugin_class.cache_clear()
AFTER_CALLS.clear()
EVENTS.clear()

def tearDown(self):
_plugin_class.cache_clear()
super().tearDown()

def test_no_plugins_keeps_response_unchanged(self):
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertFalse(AFTER_CALLS)

def test_before_hook_transforms_body(self):
seed_plugin("mutate", BODY_MUTATING_PLUGIN)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)

def test_before_hook_can_block(self):
seed_plugin("guard", BLOCKING_PLUGIN)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 403)
self.assertEqual(resp.json()["error"]["message"], "forbidden by guard")

def test_after_hook_observes_response(self):
seed_plugin("audit", AFTER_PLUGIN)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertEqual(AFTER_CALLS, [True])

def test_router_error_still_converted_by_catch(self):
seed_plugin("audit", AFTER_PLUGIN)
bad = MockConfig(
status_code=400,
response_data={"error": {"message": "upstream boom", "type": "invalid_request_error"}},
)
with self.mock_server.patch_external_api("chat/completions", bad):
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 400)
self.assertIn("upstream boom", resp.json()["error"]["message"])

def test_plugins_execute_in_order(self):
seed_plugin("first", PLUGIN_RECORD_A, order=0)
seed_plugin("second", PLUGIN_RECORD_B, order=1)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertEqual(EVENTS, ["A", "B"])

def test_plugin_execution_order_follows_order_field(self):
seed_plugin("later", PLUGIN_RECORD_A, order=1)
seed_plugin("earlier", PLUGIN_RECORD_B, order=0)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertEqual(EVENTS, ["B", "A"])

def test_blocking_plugin_short_circuits_later_plugins(self):
seed_plugin("guard", BLOCKING_PLUGIN, order=0)
seed_plugin("audit", PLUGIN_RECORD_B, order=1)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 403)
self.assertEqual(EVENTS, [], "later plugin should not run once an earlier one blocks")

def test_before_pipeline_transforms_for_next_plugin(self):
seed_plugin("set", PLUGIN_SET_STEP, order=0)
seed_plugin("read", PLUGIN_READ_STEP, order=1)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertEqual(EVENTS, ["from-pipe"])

def test_inactive_plugin_does_not_run(self):
Snippet.objects.create(
name="off", type=SnippetType.PLUGIN, active=False, code=PLUGIN_RECORD_A
)
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 200)
self.assertEqual(EVENTS, [])

def test_generic_exception_in_before_becomes_error(self):
seed_plugin("boom", PLUGIN_BOOM)
self.client.raise_request_exception = False
resp = self._send_chat_completion(self.MESSAGES)
self.assertEqual(resp.status_code, 500)
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/chat_completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -22,6 +22,7 @@
parse_body,
process_file_content,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -40,6 +41,7 @@
@check_model_availability
@normalize_reasoning_fields
@log_request
@run_plugins
@catch_router_exceptions
async def chat_completions(
request: ASGIRequest,
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/completions.py
Original file line number Diff line number Diff line change
Expand Up @@ -20,6 +20,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -36,6 +37,7 @@
@resolve_alias
@check_model_availability
@log_request
@run_plugins
@catch_router_exceptions
async def completions(
request: ASGIRequest,
Expand Down
33 changes: 33 additions & 0 deletions aqueduct/gateway/views/decorators.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,13 @@
from gateway.views.errors import error_response
from gateway.views.utils import get_response_from_cache, in_wildcard
from management.models import FileObject, Request, Token, VectorStore
from management.plugins import (
BlockedByPlugin,
PluginSnippet,
after_hook,
before_hook,
resolve_active_plugins,
)

log = logging.getLogger("aqueduct")

Expand Down Expand Up @@ -108,6 +115,7 @@ async def wrapper(request: ASGIRequest, *args: Any, **kwargs: Any) -> ViewResult
log.error("Token not found during authentication")
return unauthorized_response
kwargs["token"] = token
request.active_plugins = await sync_to_async(resolve_active_plugins)() # type: ignore[attr-defined]
return await view_func(request, *args, **kwargs)

return wrapper
Expand Down Expand Up @@ -644,6 +652,31 @@ async def wrapper(request: ASGIRequest, *args: Any, **kwargs: Any) -> ViewResult
return wrapper


def run_plugins(view_func: AsyncView) -> AsyncView:
@wraps(view_func)
async def wrapper(request: ASGIRequest, *args: Any, **kwargs: Any) -> ViewResult:
plugins: list[PluginSnippet] = getattr(request, "active_plugins", [])
if not plugins:
return await view_func(request, *args, **kwargs)

token = kwargs.get("token")
body = kwargs.get("pydantic_model")

try:
final_body = before_hook(plugins, request, token, body)
except BlockedByPlugin as e:
return error_response(e.reason, status=e.status)

if final_body is not None:
kwargs["pydantic_model"] = final_body

response = await view_func(request, *args, **kwargs)
after_hook(plugins, request, token, response)
return response

return wrapper


def catch_router_exceptions(view_func: AsyncView) -> AsyncView:
def _r(e: Exception) -> str:
s = str(e)
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/embeddings.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -37,6 +38,7 @@
@resolve_alias
@check_model_availability
@log_request
@run_plugins
@catch_router_exceptions
async def embeddings(
request: ASGIRequest,
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/image_generation.py
Original file line number Diff line number Diff line change
Expand Up @@ -18,6 +18,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -35,6 +36,7 @@
@resolve_alias
@check_model_availability
@log_request
@run_plugins
@catch_router_exceptions
async def image_generation(
request: ASGIRequest,
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/responses.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
validate_response_id,
Expand All @@ -45,6 +46,7 @@
@check_model_availability
@check_tool_availability
@log_request
@run_plugins
@catch_router_exceptions
async def create_response(
request: ASGIRequest,
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/speech.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -34,6 +35,7 @@
@ensure_usage
@check_model_availability
@log_request
@run_plugins
@catch_router_exceptions
async def speech(
request: ASGIRequest,
Expand Down
2 changes: 2 additions & 0 deletions aqueduct/gateway/views/transcriptions.py
Original file line number Diff line number Diff line change
Expand Up @@ -19,6 +19,7 @@
log_request,
parse_body,
resolve_alias,
run_plugins,
token_authenticated,
tos_accepted,
)
Expand All @@ -40,6 +41,7 @@ class TranscriptionCreateParams(RootModel): # type: ignore[type-arg]
@resolve_alias
@check_model_availability
@log_request
@run_plugins
@catch_router_exceptions
async def transcriptions(
request: ASGIRequest,
Expand Down
Loading
Loading