121 lines
3.6 KiB
Python
121 lines
3.6 KiB
Python
# retoor <retoor@molodetz.nl>
|
|||
|
|
|
||
|
|
import http.server
|
||
|
|
import json
|
||
|
|
import socket
|
||
|
|
import socketserver
|
||
|
|
import threading
|
||
|
|
|
||
|
|
import requests
|
||
|
|
|
||
|
|
from tests.conftest import BASE_URL
|
||
|
|
from tests.api.admin.gateway.index import (
|
||
|
|
JSON_gateway,
|
||
|
|
admin_session,
|
||
|
|
member_key,
|
||
|
|
_unique_gateway,
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def _fake_upstream(status_code, payload):
|
||
|
|
sock = socket.socket()
|
||
|
|
sock.bind(("127.0.0.1", 0))
|
||
|
|
port = sock.getsockname()[1]
|
||
|
|
sock.close()
|
||
|
|
body = json.dumps(payload).encode() if payload is not None else b""
|
||
|
|
seen_auth = {}
|
||
|
|
|
||
|
|
class Handler(http.server.BaseHTTPRequestHandler):
|
||
|
|
def do_GET(self):
|
||
|
|
seen_auth["path"] = self.path
|
||
|
|
seen_auth["authorization"] = self.headers.get("Authorization")
|
||
|
|
self.send_response(status_code)
|
||
|
|
self.send_header("Content-Type", "application/json")
|
||
|
|
self.end_headers()
|
||
|
|
self.wfile.write(body)
|
||
|
|
|
||
|
|
def log_message(self, *args):
|
||
|
|
pass
|
||
|
|
|
||
|
|
httpd = socketserver.TCPServer(("127.0.0.1", port), Handler)
|
||
|
|
thread = threading.Thread(target=httpd.serve_forever, daemon=True)
|
||
|
|
thread.start()
|
||
|
|
return httpd, port, seen_auth
|
||
|
|
|
||
|
|
|
||
|
|
def test_provider_models_requires_admin(seeded_db):
|
||
|
|
assert (
|
||
|
|
requests.get(
|
||
|
|
f"{BASE_URL}/admin/gateway/provider-models",
|
||
|
|
headers=JSON_gateway,
|
||
|
|
allow_redirects=False,
|
||
|
|
).status_code
|
||
|
|
== 401
|
||
|
|
)
|
||
|
|
key = member_key()
|
||
|
|
assert (
|
||
|
|
requests.get(
|
||
|
|
f"{BASE_URL}/admin/gateway/provider-models",
|
||
|
|
headers={**JSON_gateway, "X-API-KEY": key},
|
||
|
|
allow_redirects=False,
|
||
|
|
).status_code
|
||
|
|
== 403
|
||
|
|
)
|
||
|
|
|
||
|
|
|
||
|
|
def test_provider_models_returns_the_list_on_success(seeded_db):
|
||
|
|
admin = admin_session(seeded_db)
|
||
|
|
name = _unique_gateway("modelsupstream").lower()
|
||
|
|
httpd, port, seen = _fake_upstream(
|
||
|
|
200, {"data": [{"id": "vendor/a"}, {"id": "vendor/b"}]}
|
||
|
|
)
|
||
|
|
try:
|
||
|
|
admin.post(
|
||
|
|
f"{BASE_URL}/admin/gateway/providers",
|
||
|
|
json={
|
||
|
|
"name": name,
|
||
|
|
"base_url": f"http://127.0.0.1:{port}/v1/chat/completions",
|
||
|
|
"api_key": "sk-fake",
|
||
|
|
},
|
||
|
|
)
|
||
|
|
response = admin.get(
|
||
|
|
f"{BASE_URL}/admin/gateway/provider-models", params={"provider": name}
|
||
|
|
)
|
||
|
|
assert response.status_code == 200
|
||
|
|
assert response.json()["models"] == ["vendor/a", "vendor/b"]
|
||
|
|
assert seen["path"] == "/v1/models"
|
||
|
|
assert seen["authorization"] == "Bearer sk-fake"
|
||
|
|
finally:
|
||
|
|
httpd.shutdown()
|
||
|
|
admin.delete(f"{BASE_URL}/admin/gateway/providers/{name}")
|
||
|
|
|
||
|
|
|
||
|
|
def test_provider_models_404_when_upstream_has_no_model_listing(seeded_db):
|
||
|
|
admin = admin_session(seeded_db)
|
||
|
|
name = _unique_gateway("modelsupstream404").lower()
|
||
|
|
httpd, port, _ = _fake_upstream(404, {})
|
||
|
|
try:
|
||
|
|
admin.post(
|
||
|
|
f"{BASE_URL}/admin/gateway/providers",
|
||
|
|
json={
|
||
|
|
"name": name,
|
||
|
|
"base_url": f"http://127.0.0.1:{port}/v1/chat/completions",
|
||
|
|
},
|
||
|
|
)
|
||
|
|
response = admin.get(
|
||
|
|
f"{BASE_URL}/admin/gateway/provider-models", params={"provider": name}
|
||
|
|
)
|
||
|
|
assert response.status_code == 404
|
||
|
|
finally:
|
||
|
|
httpd.shutdown()
|
||
|
|
admin.delete(f"{BASE_URL}/admin/gateway/providers/{name}")
|
||
|
|
|
||
|
|
|
||
|
|
def test_provider_models_404_for_unknown_provider(seeded_db):
|
||
|
|
admin = admin_session(seeded_db)
|
||
|
|
response = admin.get(
|
||
|
|
f"{BASE_URL}/admin/gateway/provider-models",
|
||
|
|
params={"provider": _unique_gateway("ghostprovider")},
|
||
|
|
)
|
||
|
|
assert response.status_code == 404
|