use mock db for cluade code marketplace
This commit is contained in:
parent
231023c422
commit
1a9a7df437
@ -6,9 +6,12 @@ Tests:
|
||||
2. Get marketplace.json (list enabled plugins)
|
||||
"""
|
||||
|
||||
import json
|
||||
import os
|
||||
import sys
|
||||
import time
|
||||
from datetime import datetime, timezone
|
||||
from unittest.mock import AsyncMock, MagicMock
|
||||
|
||||
import pytest
|
||||
|
||||
@ -17,7 +20,6 @@ sys.path.insert(0, os.path.abspath("../.."))
|
||||
import litellm
|
||||
from litellm.proxy._types import UserAPIKeyAuth
|
||||
from litellm.proxy.proxy_server import LitellmUserRoles
|
||||
from litellm.proxy.utils import PrismaClient, ProxyLogging
|
||||
from litellm.caching.caching import DualCache
|
||||
from litellm.types.proxy.claude_code_endpoints import RegisterPluginRequest
|
||||
|
||||
@ -27,33 +29,118 @@ from litellm.proxy.anthropic_endpoints.claude_code_endpoints.claude_code_marketp
|
||||
get_marketplace,
|
||||
)
|
||||
|
||||
proxy_logging_obj = ProxyLogging(user_api_key_cache=DualCache())
|
||||
|
||||
class MockPluginRecord:
|
||||
"""Mock plugin record that mimics Prisma model behavior."""
|
||||
|
||||
def __init__(self, name, version, description, manifest_json, enabled=True, created_by=None):
|
||||
self.id = f"plugin-{name}-{int(time.time())}"
|
||||
self.name = name
|
||||
self.version = version
|
||||
self.description = description
|
||||
self.manifest_json = manifest_json
|
||||
self.files_json = "{}"
|
||||
self.enabled = enabled
|
||||
self.created_at = datetime.now(timezone.utc)
|
||||
self.updated_at = datetime.now(timezone.utc)
|
||||
self.created_by = created_by
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
def prisma_client():
|
||||
from litellm.proxy.proxy_cli import append_query_params
|
||||
def mock_prisma_client():
|
||||
"""Create a mock PrismaClient that doesn't require Prisma binaries."""
|
||||
# In-memory storage for plugins
|
||||
plugins_store = {}
|
||||
|
||||
params = {"connection_limit": 100, "pool_timeout": 60}
|
||||
database_url = os.getenv("DATABASE_URL")
|
||||
modified_url = append_query_params(database_url, params)
|
||||
os.environ["DATABASE_URL"] = modified_url
|
||||
# Create mock client
|
||||
mock_client = MagicMock()
|
||||
mock_client.proxy_logging_obj = MagicMock()
|
||||
|
||||
prisma_client = PrismaClient(
|
||||
database_url=os.environ["DATABASE_URL"], proxy_logging_obj=proxy_logging_obj
|
||||
)
|
||||
# Mock the db attribute
|
||||
mock_client.db = MagicMock()
|
||||
|
||||
litellm.proxy.proxy_server.litellm_proxy_budget_name = (
|
||||
f"litellm-proxy-budget-{time.time()}"
|
||||
)
|
||||
# Mock the plugin table with async methods
|
||||
mock_table = MagicMock()
|
||||
|
||||
return prisma_client
|
||||
async def find_unique(where):
|
||||
"""Mock find_unique - returns plugin if exists, None otherwise."""
|
||||
plugin_name = where.get("name")
|
||||
return plugins_store.get(plugin_name)
|
||||
|
||||
async def find_many(where=None):
|
||||
"""Mock find_many - returns list of plugins matching where clause."""
|
||||
if where is None or where == {}:
|
||||
return list(plugins_store.values())
|
||||
enabled = where.get("enabled")
|
||||
if enabled is not None:
|
||||
return [p for p in plugins_store.values() if p.enabled == enabled]
|
||||
return list(plugins_store.values())
|
||||
|
||||
async def create(data):
|
||||
"""Mock create - creates a new plugin."""
|
||||
plugin_name = data["name"]
|
||||
manifest = data.get("manifest_json", "{}")
|
||||
plugin = MockPluginRecord(
|
||||
name=plugin_name,
|
||||
version=data.get("version"),
|
||||
description=data.get("description"),
|
||||
manifest_json=manifest,
|
||||
enabled=data.get("enabled", True),
|
||||
created_by=data.get("created_by"),
|
||||
)
|
||||
plugins_store[plugin_name] = plugin
|
||||
return plugin
|
||||
|
||||
async def update(where, data):
|
||||
"""Mock update - updates an existing plugin."""
|
||||
plugin_name = where.get("name")
|
||||
if plugin_name not in plugins_store:
|
||||
raise ValueError(f"Plugin {plugin_name} not found")
|
||||
plugin = plugins_store[plugin_name]
|
||||
# Update fields
|
||||
if "version" in data:
|
||||
plugin.version = data["version"]
|
||||
if "description" in data:
|
||||
plugin.description = data["description"]
|
||||
if "manifest_json" in data:
|
||||
plugin.manifest_json = data["manifest_json"]
|
||||
if "enabled" in data:
|
||||
plugin.enabled = data["enabled"]
|
||||
if "updated_at" in data:
|
||||
plugin.updated_at = data["updated_at"]
|
||||
return plugin
|
||||
|
||||
async def delete(where):
|
||||
"""Mock delete - deletes a plugin."""
|
||||
plugin_name = where.get("name")
|
||||
if plugin_name in plugins_store:
|
||||
del plugins_store[plugin_name]
|
||||
return None
|
||||
|
||||
async def connect():
|
||||
"""Mock connect - no-op."""
|
||||
pass
|
||||
|
||||
# Set up async mocks
|
||||
mock_table.find_unique = AsyncMock(side_effect=find_unique)
|
||||
mock_table.find_many = AsyncMock(side_effect=find_many)
|
||||
mock_table.create = AsyncMock(side_effect=create)
|
||||
mock_table.update = AsyncMock(side_effect=update)
|
||||
mock_table.delete = AsyncMock(side_effect=delete)
|
||||
|
||||
mock_client.db.litellm_claudecodeplugintable = mock_table
|
||||
mock_client.connect = AsyncMock(side_effect=connect)
|
||||
|
||||
# Store plugins_store on the mock for cleanup if needed
|
||||
mock_client._plugins_store = plugins_store
|
||||
|
||||
return mock_client
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_register_plugin(prisma_client):
|
||||
async def test_register_plugin(mock_prisma_client):
|
||||
"""Test registering a plugin in the marketplace."""
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
@ -85,16 +172,23 @@ async def test_register_plugin(prisma_client):
|
||||
assert response["plugin"]["version"] == "1.0.0"
|
||||
assert response["plugin"]["enabled"] is True
|
||||
|
||||
# Verify the plugin was stored in the mock
|
||||
stored_plugin = await mock_prisma_client.db.litellm_claudecodeplugintable.find_unique(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
assert stored_plugin is not None
|
||||
assert stored_plugin.name == plugin_name
|
||||
|
||||
# Cleanup - delete the plugin
|
||||
await prisma_client.db.litellm_claudecodeplugintable.delete(
|
||||
await mock_prisma_client.db.litellm_claudecodeplugintable.delete(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_get_marketplace(prisma_client):
|
||||
async def test_get_marketplace(mock_prisma_client):
|
||||
"""Test getting marketplace.json with registered plugins."""
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "prisma_client", mock_prisma_client)
|
||||
setattr(litellm.proxy.proxy_server, "master_key", "sk-1234")
|
||||
|
||||
await litellm.proxy.proxy_server.prisma_client.connect()
|
||||
@ -124,7 +218,6 @@ async def test_get_marketplace(prisma_client):
|
||||
response = await get_marketplace()
|
||||
|
||||
# Response is a JSONResponse, get the body
|
||||
import json
|
||||
body = json.loads(response.body.decode())
|
||||
|
||||
assert body["name"] == "litellm"
|
||||
@ -140,6 +233,6 @@ async def test_get_marketplace(prisma_client):
|
||||
assert our_plugin["version"] == "2.0.0"
|
||||
|
||||
# Cleanup
|
||||
await prisma_client.db.litellm_claudecodeplugintable.delete(
|
||||
await mock_prisma_client.db.litellm_claudecodeplugintable.delete(
|
||||
where={"name": plugin_name}
|
||||
)
|
||||
|
||||
Loading…
Reference in New Issue
Block a user