diff --git a/.github/PULL_REQUEST_TEMPLATE.md b/.github/PULL_REQUEST_TEMPLATE.md index de50e4d9..cff9bc79 100755 --- a/.github/PULL_REQUEST_TEMPLATE.md +++ b/.github/PULL_REQUEST_TEMPLATE.md @@ -1,5 +1,14 @@ -*Issue #, if available:* +## Description -*Description of changes:* + -By submitting this pull request, I confirm that you can use, modify, copy, and redistribute this contribution, under the terms of your choice. +## Linked Issues + + + +## Checklist + +- [ ] I have performed a self-review of my code +- [ ] I have added appropriate tests +- [ ] I have updated the Defang CLI docs and/or README to reflect my changes, if necessary + \ No newline at end of file diff --git a/.github/workflows/aws-genai-cicd-suite.yml b/.github/workflows/aws-genai-cicd-suite.yml index b16c41b8..6881bd2c 100644 --- a/.github/workflows/aws-genai-cicd-suite.yml +++ b/.github/workflows/aws-genai-cicd-suite.yml @@ -51,3 +51,14 @@ jobs: echo "GitHub Token is not set" fi echo "AWS_ROLE_TO_ASSUME: ${{ vars.AWS_ROLE_TO_ASSUME_VAR }}" + + # lint the code using ruff + - name: Run Ruff linter + uses: astral-sh/ruff-action@v1 + with: + args: check ./src + + - name: Run Ruff formatter check + uses: astral-sh/ruff-action@v1 + with: + args: "format --check" diff --git a/Makefile b/Makefile index a78693f0..d337ca3b 100644 --- a/Makefile +++ b/Makefile @@ -33,3 +33,19 @@ login: ## Login to docker .PHONY: tests tests: PYTHONPATH=src pytest + +.PHONY: lint +lint: # Run pre-commit on staged/changed files + pre-commit run + +.PHONY: check +check: # Run all pre-commit hooks on all files (useful for CI or full check) + pre-commit run --all-files + +.PHONY: format +format: # Manually run ruff formatter on all files + ruff format . + +.PHONY: pre-commit-install +pre-commit-install: # Install pre-commit hooks changes + pre-commit install diff --git a/src/api/app.py b/src/api/app.py index b276f25c..295fba34 100644 --- a/src/api/app.py +++ b/src/api/app.py @@ -1,15 +1,16 @@ import logging import os -import uvicorn +import uvicorn from fastapi import FastAPI from fastapi.exceptions import RequestValidationError from fastapi.middleware.cors import CORSMiddleware from fastapi.responses import PlainTextResponse from mangum import Mangum -from api.setting import API_ROUTE_PREFIX, DESCRIPTION, SUMMARY, PROVIDER, TITLE, USE_MODEL_MAPPING, VERSION from api.modelmapper import load_model_map +from api.setting import API_ROUTE_PREFIX, DESCRIPTION, PROVIDER, SUMMARY, TITLE, USE_MODEL_MAPPING, VERSION + def is_aws(): env = os.getenv("AWS_EXECUTION_ENV") @@ -21,8 +22,9 @@ def is_aws(): return True return False + provider = PROVIDER.lower() if PROVIDER else None -if provider == None: +if provider is None: if is_aws(): provider = "aws" else: @@ -55,30 +57,36 @@ def is_aws(): if provider != "aws": from api.routers.gcp import chat, embeddings - logging.info(f"Proxy target set to: GCP") + + logging.info("Proxy target set to: GCP") app.include_router(chat.router, prefix=API_ROUTE_PREFIX) app.include_router(embeddings.router, prefix=API_ROUTE_PREFIX) else: from api.routers import chat, embeddings, model + logging.info("No proxy target set. Using internal routers.") app.include_router(model.router, prefix=API_ROUTE_PREFIX) app.include_router(chat.router, prefix=API_ROUTE_PREFIX) app.include_router(embeddings.router, prefix=API_ROUTE_PREFIX) + @app.get("/", include_in_schema=False) async def root(): """Root endpoint for the API""" return {"status": "OK"} + @app.get("/health") async def health(): """For health check if needed""" return {"status": "OK"} + @app.exception_handler(RequestValidationError) async def validation_exception_handler(request, exc): return PlainTextResponse(str(exc), status_code=400) + handler = Mangum(app) if __name__ == "__main__": diff --git a/src/api/auth.py b/src/api/auth.py index 09aa3fb1..e927bad0 100644 --- a/src/api/auth.py +++ b/src/api/auth.py @@ -28,7 +28,7 @@ raise RuntimeError("Unable to retrieve API KEY, please ensure the secret ARN is correct") except KeyError: raise RuntimeError('Please ensure the secret contains a "api_key" field') -elif api_key_env != None: +elif api_key_env is not None: api_key = api_key_env else: # For local use only. diff --git a/src/api/gcp/credentials/metadata.py b/src/api/gcp/credentials/metadata.py index 80a57bb9..59c42576 100644 --- a/src/api/gcp/credentials/metadata.py +++ b/src/api/gcp/credentials/metadata.py @@ -1,17 +1,17 @@ import logging -import requests +import requests from google.auth import default from google.auth.transport.requests import Request as AuthRequest -from api.setting import GOOGLE_CLOUD_PROJECT, GCP_REGION - +from api.setting import GCP_REGION, GOOGLE_CLOUD_PROJECT # GCP credentials and project details credentials = None project_id = None location = None + def get_gcp_project_details(): from google.auth import default @@ -29,17 +29,19 @@ def get_gcp_project_details(): zone = requests.get( "http://metadata.google.internal/computeMetadata/v1/instance/zone", headers={"Metadata-Flavor": "Google"}, - timeout=1 + timeout=1, ).text location = zone.split("/")[-1].rsplit("-", 1)[0] except Exception: - logging.warning(f"Error: Failed to get project and location from metadata server. Using local settings.") + logging.warning("Error: Failed to get project and location from metadata server. Using local settings.") return credentials, project_id, location + credentials, project_id, location = get_gcp_project_details() + # Utility: get service account access token def get_access_token(): credentials, _ = default(scopes=["https://www.googleapis.com/auth/cloud-platform"]) diff --git a/src/api/modelmapper.py b/src/api/modelmapper.py index b550affd..fb3a04a3 100644 --- a/src/api/modelmapper.py +++ b/src/api/modelmapper.py @@ -1,9 +1,10 @@ -import os import json +import os from pathlib import Path _model_map = None + def load_model_map(): global _model_map BASE_DIR = os.path.dirname(os.path.abspath(__file__)) @@ -11,6 +12,7 @@ def load_model_map(): with open(modelmap_path, "r") as f: _model_map = json.load(f) + def get_model(provider, model, fallback_model): provider = provider.lower() if model is None or model == "": @@ -19,4 +21,3 @@ def get_model(provider, model, fallback_model): available_models = _model_map.get(provider, {}) return available_models.get(model, model) - diff --git a/src/api/models/bedrock.py b/src/api/models/bedrock.py index a4d0f892..2fe250d5 100644 --- a/src/api/models/bedrock.py +++ b/src/api/models/bedrock.py @@ -14,6 +14,7 @@ from fastapi import HTTPException from starlette.concurrency import run_in_threadpool +from api.modelmapper import get_model from api.models.base import BaseChatModel, BaseEmbeddingsModel from api.schema import ( AssistantMessage, @@ -39,7 +40,6 @@ UserMessage, ) from api.setting import AWS_REGION, DEBUG, DEFAULT_MODEL, ENABLE_CROSS_REGION_INFERENCE -from api.modelmapper import get_model logger = logging.getLogger(__name__) @@ -145,16 +145,15 @@ def validate(self, chat_request: ChatRequest): if DEBUG: logger.debug("Bedrock validate " + chat_request.model + " list: " + json.dumps(bedrock_model_list)) logger.debug(f"Checking model: {repr(chat_request.model)}") - logger.debug(f"Available keys include: {repr('anthropic.claude-3-5-sonnet-20241022-v2:0') in bedrock_model_list}") + logger.debug( + f"Available keys include: {repr('anthropic.claude-3-5-sonnet-20241022-v2:0') in bedrock_model_list}" + ) # check if model is supported if chat_request.model not in bedrock_model_list.keys(): if DEBUG: logger.debug(f"Bedrock list: {list(bedrock_model_list.keys())}") - error = ( - f"Unsupported model '{chat_request.model}'. " - f"list of known models: {bedrock_model_list.keys()}" - ) + error = f"Unsupported model '{chat_request.model}'. list of known models: {bedrock_model_list.keys()}" logger.error(error) if error: diff --git a/src/api/routers/chat.py b/src/api/routers/chat.py index b74d00c0..4a5a709c 100644 --- a/src/api/routers/chat.py +++ b/src/api/routers/chat.py @@ -4,10 +4,9 @@ from fastapi.responses import StreamingResponse from api.auth import api_key_auth +from api.modelmapper import get_model from api.models.bedrock import BedrockModel from api.schema import ChatRequest, ChatResponse, ChatStreamResponse, Error -from api.modelmapper import get_model - from api.setting import DEFAULT_MODEL, USE_MODEL_MAPPING router = APIRouter( @@ -36,10 +35,10 @@ async def chat_completions( ), ], ): - if chat_request.model != None and chat_request.model.lower().startswith("gpt-"): + if chat_request.model is not None and chat_request.model.lower().startswith("gpt-"): chat_request.model = DEFAULT_MODEL - # replace with mapped model name + # replace with mapped model name if USE_MODEL_MAPPING: req_model = chat_request.model req_model = get_model("aws", req_model, "chat-default") diff --git a/src/api/routers/embeddings.py b/src/api/routers/embeddings.py index eca56e54..28fe2581 100644 --- a/src/api/routers/embeddings.py +++ b/src/api/routers/embeddings.py @@ -3,10 +3,10 @@ from fastapi import APIRouter, Body, Depends from api.auth import api_key_auth +from api.modelmapper import get_model from api.models.bedrock import get_embeddings_model from api.schema import EmbeddingsRequest, EmbeddingsResponse from api.setting import DEFAULT_EMBEDDING_MODEL -from api.modelmapper import get_model router = APIRouter( prefix="/embeddings", @@ -28,7 +28,7 @@ async def embeddings( ), ], ): - if embeddings_request.model != None and embeddings_request.model.lower().startswith("text-embedding-"): + if embeddings_request.model is not None and embeddings_request.model.lower().startswith("text-embedding-"): embeddings_request.model = DEFAULT_EMBEDDING_MODEL # Exception will be raised if model not supported. embeddings_request.model = get_model("aws", embeddings_request.model, "embedding-default") diff --git a/src/api/routers/gcp/chat.py b/src/api/routers/gcp/chat.py index 26c69133..12f0edbe 100644 --- a/src/api/routers/gcp/chat.py +++ b/src/api/routers/gcp/chat.py @@ -1,21 +1,21 @@ -import httpx import json import logging import os - +from contextlib import asynccontextmanager from typing import AsyncGenerator + +import httpx from fastapi import APIRouter, Depends, Request, Response from fastapi.responses import StreamingResponse -from contextlib import asynccontextmanager -from api.setting import API_ROUTE_PREFIX, USE_MODEL_MAPPING from google.auth import default from google.auth.transport.requests import Request as AuthRequest from api.auth import api_key_auth +from api.gcp.credentials.metadata import get_access_token, location, project_id from api.modelmapper import get_model -from api.gcp.credentials.metadata import get_access_token, project_id, location +from api.routers.gcp.stream_transformers import handle_data_line, sse_chunk, sse_done from api.schema import ChatResponse, ChatStreamResponse, Error -from api.routers.gcp.stream_transformers import handle_data_line, sse_done, sse_chunk +from api.setting import API_ROUTE_PREFIX, USE_MODEL_MAPPING known_chat_models = [ "publishers/mistral-ai/models/mistral-7b-instruct-v0.3", @@ -40,6 +40,7 @@ responses={404: {"description": "Not found"}}, ) + def get_proxy_target(model, path, stream): """ Check if the environment variable is set to use GCP. @@ -52,6 +53,7 @@ def get_proxy_target(model, path, stream): endPointSuffix = "streamRawPredict" if stream else "rawPredict" return f"https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/{model}:{endPointSuffix}" + def get_headers(model, request, path, stream): target_url = None path_no_prefix = f"/{path.lstrip('/')}".removeprefix(API_ROUTE_PREFIX) @@ -59,7 +61,8 @@ def get_headers(model, request, path, stream): # remove hop-by-hop headers headers = { - k: v for k, v in request.headers.items() + k: v + for k, v in request.headers.items() if k.lower() not in {"host", "content-length", "accept-encoding", "connection", "authorization"} } @@ -68,6 +71,7 @@ def get_headers(model, request, path, stream): headers["Authorization"] = f"Bearer {access_token}" return target_url, headers + def _parse_system_prompts(openai_messages) -> str: system_prompts = "" for message in openai_messages: @@ -79,59 +83,54 @@ def _parse_system_prompts(openai_messages) -> str: return system_prompts + def to_vertex_anthropic(openai_messages): message = [] for m in openai_messages["messages"]: if m["role"] == "system": continue - message.append({ - "role": m["role"], - "content": [{"type": "text", "text": m["content"]}] - }) + message.append({"role": m["role"], "content": [{"type": "text", "text": m["content"]}]}) system_prompts = _parse_system_prompts(openai_messages["messages"]) - return { - "anthropic_version": "vertex-2023-10-16", - "max_tokens": 256, - "system": system_prompts, - "messages": message - } + return {"anthropic_version": "vertex-2023-10-16", "max_tokens": 256, "system": system_prompts, "messages": message} + def from_anthropic_to_openai_response(msg, model): msg_json = json.loads(msg) - return json.dumps({ - "id": msg_json["id"], - "object": "chat.completion", - "model": model, - "choices": [ - { - "index": 0, - "message": { - "role": msg_json["role"], - "content": "".join( - part["text"] for part in msg_json["content"] - if part["type"] == "text" - ) - }, - "finish_reason": msg_json.get("stop_reason", "stop") - } - ], - "usage": msg_json.get("usage", {}) - }) + return json.dumps( + { + "id": msg_json["id"], + "object": "chat.completion", + "model": model, + "choices": [ + { + "index": 0, + "message": { + "role": msg_json["role"], + "content": "".join(part["text"] for part in msg_json["content"] if part["type"] == "text"), + }, + "finish_reason": msg_json.get("stop_reason", "stop"), + } + ], + "usage": msg_json.get("usage", {}), + } + ) + def get_chat_completion_model_name(model_alias): if model_alias.startswith("publishers/google/"): return f"google/{model_alias.split('/')[-1]}" - return model_alias.split('/')[-1] + return model_alias.split("/")[-1] + def transform_vertex_chunk_to_openai(chunk_line: str, index: int = 0) -> str: if not chunk_line.startswith("data: "): return "" # skip irrelevant lines like empty keepalives try: - payload = json.loads(chunk_line[len("data: "):]) + payload = json.loads(chunk_line[len("data: ") :]) parts = payload.get("candidates", [])[0].get("content", {}).get("parts", []) if not parts: return "" @@ -139,19 +138,14 @@ def transform_vertex_chunk_to_openai(chunk_line: str, index: int = 0) -> str: except Exception: return "" - transformed = { - "choices": [ - { - "delta": {"content": text}, - "index": index, - "finish_reason": None - } - ] - } + transformed = {"choices": [{"delta": {"content": text}, "index": index, "finish_reason": None}]} return f"data: {json.dumps(transformed)}\n" -async def stream_generator(target_url: str, request_headers: dict, content_json: dict, model_alias: str) -> AsyncGenerator[str, None]: + +async def stream_generator( + target_url: str, request_headers: dict, content_json: dict, model_alias: str +) -> AsyncGenerator[str, None]: async with httpx.AsyncClient(timeout=None) as client: async with client.stream( "POST", @@ -159,7 +153,6 @@ async def stream_generator(target_url: str, request_headers: dict, content_json: headers=request_headers, json=content_json, ) as response: - logging.debug(f"Received response with status code: {response.status_code}") logging.debug(f"Response headers: {response.headers}") @@ -173,7 +166,7 @@ async def stream_generator(target_url: str, request_headers: dict, content_json: break if line.startswith("data: "): - raw_json = line[len("data: "):].strip() + raw_json = line[len("data: ") :].strip() else: raw_json = line.strip() async for chunk in handle_data_line(raw_json, model_alias): @@ -181,6 +174,7 @@ async def stream_generator(target_url: str, request_headers: dict, content_json: yield chunk yield sse_done() + @router.post( "/completions", response_model=ChatResponse | ChatStreamResponse | Error, response_model_exclude_unset=True ) @@ -209,7 +203,9 @@ async def handle_proxy(request: Request): logging.debug(f"Request headers: {request_headers}") logging.debug(f"Request content: {content_json}") if is_streaming: - return StreamingResponse(stream_generator(target_url, request_headers, content_json, model_alias), media_type="text/event-stream") + return StreamingResponse( + stream_generator(target_url, request_headers, content_json, model_alias), media_type="text/event-stream" + ) async with httpx.AsyncClient() as client: response = await client.request( @@ -232,7 +228,8 @@ async def handle_proxy(request: Request): # remove hop-by-hop headers response_headers = { - k: v for k, v in response.headers.items() + k: v + for k, v in response.headers.items() if k.lower() not in {"content-encoding", "transfer-encoding", "connection"} } diff --git a/src/api/routers/gcp/embeddings.py b/src/api/routers/gcp/embeddings.py index 930430b8..306b309f 100644 --- a/src/api/routers/gcp/embeddings.py +++ b/src/api/routers/gcp/embeddings.py @@ -1,20 +1,22 @@ -import httpx import json import logging import os +import httpx from fastapi import APIRouter, Depends, Request, Response + from api.auth import api_key_auth +from api.gcp.credentials.metadata import get_access_token, location, project_id +from api.modelmapper import get_model from api.schema import EmbeddingsResponse from api.setting import API_ROUTE_PREFIX -from api.modelmapper import get_model -from api.gcp.credentials.metadata import get_access_token, project_id, location - router = APIRouter( prefix="/embeddings", dependencies=[Depends(api_key_auth)], ) + + def get_proxy_target(model, path): """ Check if the environment variable is set to use GCP. @@ -24,13 +26,15 @@ def get_proxy_target(model, path): else: return f"https://{location}-aiplatform.googleapis.com/v1/projects/{project_id}/locations/{location}/{model}:predict" + def get_header(model, request, path): path_no_prefix = f"/{path.lstrip('/')}".removeprefix(API_ROUTE_PREFIX) target_url = get_proxy_target(model, path_no_prefix) # remove hop-by-hop headers headers = { - k: v for k, v in request.headers.items() + k: v + for k, v in request.headers.items() if k.lower() not in {"host", "content-length", "accept-encoding", "connection", "authorization"} } @@ -39,6 +43,7 @@ def get_header(model, request, path): headers["Authorization"] = f"Bearer {access_token}" return target_url, headers + def to_vertex_embeddings(request): """ Convert OpenAI-style embeddings request to Vertex AI format. @@ -46,17 +51,15 @@ def to_vertex_embeddings(request): inputs = request.get("input", []) if not isinstance(inputs, list): inputs = [inputs] - return { - "instances": [{"content": str(content)} for content in inputs] - } + return {"instances": [{"content": str(content)} for content in inputs]} + def to_openai_response(embedding_content, model): """ Convert Vertex AI embeddings response to OpenAI format. """ total_tokens = sum( - item["embeddings"]["statistics"]["token_count"] - for item in embedding_content.get("predictions", []) + item["embeddings"]["statistics"]["token_count"] for item in embedding_content.get("predictions", []) ) return { @@ -65,15 +68,15 @@ def to_openai_response(embedding_content, model): "embedding": item["embeddings"]["values"], "index": idx, "object": "embedding", - } for idx, item in enumerate(embedding_content.get("predictions", [])) + } + for idx, item in enumerate(embedding_content.get("predictions", [])) ], "model": model, "object": "list", - "usage": { - "total_tokens": total_tokens - } + "usage": {"total_tokens": total_tokens}, } + @router.post("/{path:path}", response_model=EmbeddingsResponse) async def handle_proxy(request: Request, path: str): try: @@ -103,7 +106,8 @@ async def handle_proxy(request: Request, path: str): # remove hop-by-hop headers response_headers = { - k: v for k, v in response.headers.items() + k: v + for k, v in response.headers.items() if k.lower() not in {"content-encoding", "transfer-encoding", "connection"} } diff --git a/src/api/routers/gcp/stream_transformers.py b/src/api/routers/gcp/stream_transformers.py index 780a29a5..cc0e16d8 100644 --- a/src/api/routers/gcp/stream_transformers.py +++ b/src/api/routers/gcp/stream_transformers.py @@ -1,61 +1,54 @@ import json -import time import logging +import time from typing import AsyncGenerator + def sse_chunk(payload: str) -> str: return f"data: {payload}\n\n" + def sse_done() -> str: return "data: [DONE]\n\n" + def generate_openai_id() -> str: return f"chatcmpl-ts{int(time.time() * 1000)}" + def transform_claude(data: dict): if data["type"] == "content_block_delta": - return json.dumps({ - "id": generate_openai_id(), - "object": "chat.completion.chunk", - "choices": [ - { - "delta": { "content": data["delta"]["text"] }, - "index": 0, - "finish_reason": None - } - ] - }) + return json.dumps( + { + "id": generate_openai_id(), + "object": "chat.completion.chunk", + "choices": [{"delta": {"content": data["delta"]["text"]}, "index": 0, "finish_reason": None}], + } + ) if data["type"] == "message_delta": - return json.dumps({ - "choices": [ - { - "delta": {}, - "index": 0, - "finish_reason": data["delta"]["stop_reason"] - } - ] - }) + return json.dumps({"choices": [{"delta": {}, "index": 0, "finish_reason": data["delta"]["stop_reason"]}]}) if data["type"] == "message": - return json.dumps({ - "id": generate_openai_id(), - "object": "chat.completion", - "model": data["model"], - "choices": [ - { - "index": 0, - "delta": { - "content": data["content"][0]["text"] - }, - "finish_reason": data.get("stop_reason", "stop") - } - ], - "usage": {} - }) + return json.dumps( + { + "id": generate_openai_id(), + "object": "chat.completion", + "model": data["model"], + "choices": [ + { + "index": 0, + "delta": {"content": data["content"][0]["text"]}, + "finish_reason": data.get("stop_reason", "stop"), + } + ], + "usage": {}, + } + ) logging.warning(f"Unknown data type: {data['type']}") + async def handle_data_line(raw_json: str, model: str) -> AsyncGenerator[str, None]: try: data = json.loads(raw_json) diff --git a/src/api/routers/gcp/test_chat.py b/src/api/routers/gcp/test_chat.py index c584a7bc..c012f455 100644 --- a/src/api/routers/gcp/test_chat.py +++ b/src/api/routers/gcp/test_chat.py @@ -1,17 +1,19 @@ -import pytest import json -from unittest.mock import patch, MagicMock -from starlette.datastructures import Headers, QueryParams +from unittest.mock import MagicMock, patch + +import pytest from fastapi import Response +from starlette.datastructures import Headers, QueryParams import api.routers.gcp.chat as chat + @pytest.fixture def dummy_request(): class DummyRequest: def __init__(self, headers=None, body=None, method="POST", query_params=None): self.headers = Headers(headers or {}) - self._body = body or b'{}' + self._body = body or b"{}" self.method = method self.query_params = QueryParams(query_params or {}) @@ -20,17 +22,18 @@ async def body(self): return DummyRequest + def test_to_vertex_anthropic(): openai_messages = { "messages": [ {"role": "system", "content": "You are a helpful assistant."}, {"role": "user", "content": "Hello!"}, - {"role": "assistant", "content": "Hi there!"} + {"role": "assistant", "content": "Hi there!"}, ] } result = chat.to_vertex_anthropic(openai_messages) assert result["anthropic_version"] == "vertex-2023-10-16" - assert result["system"] == 'You are a helpful assistant.\n' + assert result["system"] == "You are a helpful assistant.\n" assert result["max_tokens"] == 256 assert isinstance(result["messages"], list) assert result["messages"][0]["role"] == "user" @@ -38,14 +41,17 @@ def test_to_vertex_anthropic(): assert result["messages"][1]["role"] == "assistant" assert result["messages"][1]["content"][0]["text"] == "Hi there!" + def test_from_anthropic_to_openai_response(): - msg = json.dumps({ - "id": "abc123", - "role": "assistant", - "content": [{"type": "text", "text": "Hello!"}, {"type": "text", "text": "Bye!"}], - "stop_reason": "stop", - "usage": {"prompt_tokens": 5, "completion_tokens": 2} - }) + msg = json.dumps( + { + "id": "abc123", + "role": "assistant", + "content": [{"type": "text", "text": "Hello!"}, {"type": "text", "text": "Bye!"}], + "stop_reason": "stop", + "usage": {"prompt_tokens": 5, "completion_tokens": 2}, + } + ) result = json.loads(chat.from_anthropic_to_openai_response(msg, "default")) assert result["id"] == "abc123" assert result["object"] == "chat.completion" @@ -54,11 +60,13 @@ def test_from_anthropic_to_openai_response(): assert result["choices"][0]["finish_reason"] == "stop" assert result["usage"]["prompt_tokens"] == 5 + def test_get_proxy_target_env(monkeypatch): monkeypatch.setenv("PROXY_TARGET", "https://custom-proxy") result = chat.get_proxy_target("any-model", "/v1/chat/completions", False) assert result == "https://custom-proxy" + def test_get_proxy_target_known_chat(monkeypatch): monkeypatch.delenv("PROXY_TARGET", raising=False) model = chat.known_chat_models[0] @@ -66,6 +74,7 @@ def test_get_proxy_target_known_chat(monkeypatch): result = chat.get_proxy_target(model, path, False) assert "endpoints/openapi/chat/completions" in result + def test_get_proxy_target_raw_predict(monkeypatch): monkeypatch.delenv("PROXY_TARGET", raising=False) model = "unknown-model" @@ -73,6 +82,7 @@ def test_get_proxy_target_raw_predict(monkeypatch): result = chat.get_proxy_target(model, path, False) assert ":rawPredict" in result + def test_get_proxy_target_stream_raw_predict(monkeypatch): monkeypatch.delenv("PROXY_TARGET", raising=False) model = "unknown-model" @@ -80,16 +90,19 @@ def test_get_proxy_target_stream_raw_predict(monkeypatch): result = chat.get_proxy_target(model, path, True) assert ":streamRawPredict" in result + @patch("api.routers.gcp.chat.get_access_token", return_value="dummy-token") def test_get_header_removes_hop_headers(mock_token, dummy_request): - req = dummy_request(headers={ - "Host": "example.com", - "Content-Length": "123", - "Accept-Encoding": "gzip", - "Connection": "keep-alive", - "Authorization": "Bearer old", - "X-Custom": "foo" - }) + req = dummy_request( + headers={ + "Host": "example.com", + "Content-Length": "123", + "Accept-Encoding": "gzip", + "Connection": "keep-alive", + "Authorization": "Bearer old", + "X-Custom": "foo", + } + ) model = "test-model" path = "/v1/chat/completions" with patch("api.routers.gcp.chat.get_proxy_target", return_value="http://target"): @@ -103,6 +116,7 @@ def test_get_header_removes_hop_headers(mock_token, dummy_request): assert header["Authorization"] == "Bearer dummy-token" assert header["x-custom"] == "foo" + @pytest.mark.asyncio @patch("api.routers.gcp.chat.httpx.AsyncClient") @patch("api.routers.gcp.chat.get_headers") @@ -123,13 +137,12 @@ async def test_handle_proxy_basic(mock_get_model, mock_get_headers, mock_async_c assert b"hi" in result.body assert result.headers["content-type"] == "application/json" + @pytest.mark.asyncio @patch("api.routers.gcp.chat.httpx.AsyncClient") @patch("api.routers.gcp.chat.get_headers") @patch("api.routers.gcp.chat.get_model", return_value="test-model") -async def test_handle_proxy_known_chat_model( - mock_get_model, mock_get_headers, mock_async_client, dummy_request -): +async def test_handle_proxy_known_chat_model(mock_get_model, mock_get_headers, mock_async_client, dummy_request): req = dummy_request(body=json.dumps({"model": "foo"}).encode()) mock_get_headers.return_value = ("http://target", {"Authorization": "Bearer token"}) mock_response = MagicMock() @@ -148,24 +161,25 @@ async def test_handle_proxy_known_chat_model( assert b"hi" in result.body assert result.headers["content-type"] == "application/json" + @pytest.mark.asyncio @patch("api.routers.gcp.chat.httpx.AsyncClient") @patch("api.routers.gcp.chat.get_headers") @patch("api.routers.gcp.chat.get_model", return_value="anthropic-model") -async def test_handle_proxy_anthropic_conversion( - mock_get_model, mock_get_headers, mock_async_client, dummy_request -): +async def test_handle_proxy_anthropic_conversion(mock_get_model, mock_get_headers, mock_async_client, dummy_request): req = dummy_request(body=json.dumps({"model": "foo", "messages": [{"role": "user", "content": "hi"}]}).encode()) mock_get_headers.return_value = ("http://target", {"Authorization": "Bearer token"}) mock_response = MagicMock() # Simulate anthropic response - anthropic_resp = json.dumps({ - "id": "abc123", - "role": "assistant", - "content": [{"type": "text", "text": "Hello!"}], - "stop_reason": "stop", - "usage": {"prompt_tokens": 5, "completion_tokens": 2} - }).encode() + anthropic_resp = json.dumps( + { + "id": "abc123", + "role": "assistant", + "content": [{"type": "text", "text": "Hello!"}], + "stop_reason": "stop", + "usage": {"prompt_tokens": 5, "completion_tokens": 2}, + } + ).encode() mock_response.content = anthropic_resp mock_response.status_code = 200 mock_response.headers = {"content-type": "application/json"} @@ -181,13 +195,12 @@ async def test_handle_proxy_anthropic_conversion( assert data["object"] == "chat.completion" assert data["choices"][0]["message"]["content"] == "Hello!" + @pytest.mark.asyncio @patch("api.routers.gcp.chat.httpx.AsyncClient", side_effect=Exception("network error")) @patch("api.routers.gcp.chat.get_headers") @patch("api.routers.gcp.chat.get_model", return_value="test-model") -async def test_handle_proxy_httpx_exception( - mock_get_model, mock_get_headers, mock_async_client, dummy_request -): +async def test_handle_proxy_httpx_exception(mock_get_model, mock_get_headers, mock_async_client, dummy_request): req = dummy_request(body=json.dumps({"model": "foo"}).encode()) mock_get_headers.return_value = ("http://target", {"Authorization": "Bearer token"}) chat.USE_MODEL_MAPPING = True @@ -205,6 +218,7 @@ async def test_handle_proxy_httpx_exception( # Assert that the response body contains the expected error message assert b"Upstream request failed" in result.body + def test_get_chat_completion_model_name_known_chat_model(): # Pick a known chat model from the list model_alias = "publishers/google/models/gemini-2.0-flash-lite-001" @@ -219,6 +233,7 @@ def test_get_chat_completion_model_name_known_chat_model(): # Should remove 'publishers/' and 'models/' from the string assert result == "google/gemini-2.0-flash-lite-001" + def test_get_chat_completion_model_name_unknown_model(): model_alias = "some-other-model" # Ensure it's not in known_chat_models @@ -227,4 +242,3 @@ def test_get_chat_completion_model_name_unknown_model(): result = chat.get_chat_completion_model_name(model_alias) # Should return the input unchanged assert result == model_alias - diff --git a/src/api/routers/gcp/test_embeddings.py b/src/api/routers/gcp/test_embeddings.py index d6b0548c..3239b13a 100644 --- a/src/api/routers/gcp/test_embeddings.py +++ b/src/api/routers/gcp/test_embeddings.py @@ -1,129 +1,85 @@ -from api.routers.gcp.embeddings import to_vertex_embeddings -from api.routers.gcp.embeddings import to_openai_response +from api.routers.gcp.embeddings import to_openai_response, to_vertex_embeddings + def test_to_vertex_embeddings_with_string_input(): request = {"input": "hello world"} - expected = { - "instances": [ - {"content": "hello world"} - ] - } + expected = {"instances": [{"content": "hello world"}]} assert to_vertex_embeddings(request) == expected + def test_to_vertex_embeddings_with_list_of_strings(): request = {"input": ["foo", "bar"]} - expected = { - "instances": [ - {"content": "foo"}, - {"content": "bar"} - ] - } + expected = {"instances": [{"content": "foo"}, {"content": "bar"}]} assert to_vertex_embeddings(request) == expected + def test_to_vertex_embeddings_with_list_of_numbers(): request = {"input": [1, 2, 3]} - expected = { - "instances": [ - {"content": "1"}, - {"content": "2"}, - {"content": "3"} - ] - } + expected = {"instances": [{"content": "1"}, {"content": "2"}, {"content": "3"}]} assert to_vertex_embeddings(request) == expected + def test_to_vertex_embeddings_with_empty_list(): request = {"input": []} - expected = { - "instances": [] - } + expected = {"instances": []} assert to_vertex_embeddings(request) == expected + def test_to_vertex_embeddings_with_missing_input_key(): request = {} - expected = { - "instances": [] - } + expected = {"instances": []} assert to_vertex_embeddings(request) == expected + def test_to_openai_response_with_multiple_predictions(): embedding_content = { "predictions": [ {"embeddings": {"values": [0.1, 0.2, 0.3], "statistics": {"token_count": 10}}}, - {"embeddings": {"values": [0.4, 0.5, 0.6], "statistics": {"token_count": 20}}} + {"embeddings": {"values": [0.4, 0.5, 0.6], "statistics": {"token_count": 20}}}, ] } model = "test-model" expected = { "data": [ {"embedding": [0.1, 0.2, 0.3], "index": 0, "object": "embedding"}, - {"embedding": [0.4, 0.5, 0.6], "index": 1, "object": "embedding"} + {"embedding": [0.4, 0.5, 0.6], "index": 1, "object": "embedding"}, ], "model": "test-model", "object": "list", - "usage": { - "total_tokens": 30 - } + "usage": {"total_tokens": 30}, } assert to_openai_response(embedding_content, model) == expected + def test_to_openai_response_with_empty_predictions(): - embedding_content = { - "predictions": [] - } + embedding_content = {"predictions": []} model = "empty-model" - expected = { - "data": [], - "model": "empty-model", - "object": "list", - "usage": { - "total_tokens": 0 - } - } + expected = {"data": [], "model": "empty-model", "object": "list", "usage": {"total_tokens": 0}} assert to_openai_response(embedding_content, model) == expected + def test_to_openai_response_with_missing_predictions_key(): embedding_content = {} model = "default-model" - expected = { - "data": [], - "model": "default-model", - "object": "list", - "usage": { - "total_tokens": 0 - } - } + expected = {"data": [], "model": "default-model", "object": "list", "usage": {"total_tokens": 0}} assert to_openai_response(embedding_content, model) == expected + def test_to_openai_response_with_single_prediction(): - embedding_content = { - "predictions": [ - { - "embeddings": { - "values": [1.0, 2.0, 3.0], "statistics": {"token_count": 5} - } - } - ] - } + embedding_content = {"predictions": [{"embeddings": {"values": [1.0, 2.0, 3.0], "statistics": {"token_count": 5}}}]} model = "single-model" expected = { - "data": [ - {"embedding": [1.0, 2.0, 3.0], "index": 0, "object": "embedding"} - ], + "data": [{"embedding": [1.0, 2.0, 3.0], "index": 0, "object": "embedding"}], "model": "single-model", "object": "list", - "usage": { - "total_tokens": 5 - } + "usage": {"total_tokens": 5}, } assert to_openai_response(embedding_content, model) == expected + def test_to_openai_response_with_nonstandard_embedding_key(): # This should raise a KeyError if "embeddings" or "values" is missing - embedding_content = { - "predictions": [ - {"not_embeddings": {"values": [1, 2, 3]}} - ] - } + embedding_content = {"predictions": [{"not_embeddings": {"values": [1, 2, 3]}}]} model = "bad-model" try: to_openai_response(embedding_content, model) diff --git a/src/api/routers/gcp/test_stream_transformers.py b/src/api/routers/gcp/test_stream_transformers.py index e40127d3..c36e97bf 100644 --- a/src/api/routers/gcp/test_stream_transformers.py +++ b/src/api/routers/gcp/test_stream_transformers.py @@ -1,24 +1,27 @@ -import pytest +import asyncio import json import re -import asyncio +import pytest from stream_transformers import ( - sse_chunk, - sse_done, generate_openai_id, handle_data_line, + sse_chunk, + sse_done, ) + def test_sse_chunk(): payload = '{"foo": "bar"}' assert sse_chunk(payload) == f"data: {payload}\n\n" + def test_generate_openai_id(): id1 = generate_openai_id() assert id1.startswith("chatcmpl-ts") assert re.match(r"chatcmpl-ts\d{13}", id1) + @pytest.mark.asyncio async def test_handle_data_line_content_block_delta(): data = { @@ -27,12 +30,13 @@ async def test_handle_data_line_content_block_delta(): } gen = handle_data_line(json.dumps(data), "test-model") chunk = await anext(gen) - obj = json.loads(chunk[len("data: "):-2]) + obj = json.loads(chunk[len("data: ") : -2]) assert obj["object"] == "chat.completion.chunk" assert obj["choices"][0]["delta"]["content"] == "Hello!" with pytest.raises(StopAsyncIteration): await anext(gen) + @pytest.mark.asyncio async def test_handle_data_line_message_delta_with_stop_reason(): data = { @@ -43,7 +47,7 @@ async def test_handle_data_line_message_delta_with_stop_reason(): # chunk 1 chunk = await anext(gen) - obj = json.loads(chunk[len("data: "):-2]) + obj = json.loads(chunk[len("data: ") : -2]) assert obj["choices"][0]["finish_reason"] == "stop" # chunk 2 @@ -52,6 +56,7 @@ async def test_handle_data_line_message_delta_with_stop_reason(): with pytest.raises(StopAsyncIteration): await anext(gen) + @pytest.mark.asyncio async def test_handle_data_line_message_stop(): data = {"type": "message_stop"} @@ -61,22 +66,20 @@ async def test_handle_data_line_message_stop(): with pytest.raises(StopAsyncIteration): await anext(gen) + @pytest.mark.asyncio async def test_handle_data_line_claude_content_block_delta(): - data = { - "type": "content_block_delta", - "delta": {"text": "Hi!"}, - "model": "claude-2" - } + data = {"type": "content_block_delta", "delta": {"text": "Hi!"}, "model": "claude-2"} raw_json = json.dumps(data) gen = handle_data_line(raw_json, "override-model") chunk = await anext(gen) - obj = json.loads(chunk[len("data: "):-2]) + obj = json.loads(chunk[len("data: ") : -2]) assert obj["choices"][0]["delta"]["content"] == "Hi!" - assert "model" not in obj # model should not be in delta blocks + assert "model" not in obj # model should not be in delta blocks with pytest.raises(StopAsyncIteration): await anext(gen) + @pytest.mark.asyncio async def test_handle_data_line_gemini_or_openai(): data = { @@ -87,12 +90,13 @@ async def test_handle_data_line_gemini_or_openai(): raw_json = json.dumps(data) gen = handle_data_line(raw_json, "override-model") chunk = await anext(gen) - obj = json.loads(chunk[len("data: "):-2]) + obj = json.loads(chunk[len("data: ") : -2]) assert obj["id"] == "abc" assert obj["model"] == "override-model" with pytest.raises(StopAsyncIteration): await anext(gen) + @pytest.mark.asyncio async def test_handle_data_line_invalid_json(): raw_json = "not a json" diff --git a/src/api/test_modelmapper.py b/src/api/test_modelmapper.py index 9bd13f1b..2e8c5037 100644 --- a/src/api/test_modelmapper.py +++ b/src/api/test_modelmapper.py @@ -1,25 +1,24 @@ import unittest -from unittest.mock import patch, mock_open +from unittest.mock import mock_open, patch + from api.modelmapper import get_model, load_model_map -@patch("api.modelmapper._model_map", { - "provider1": { - "model1": "mapped_model1", - "model2": "mapped_model2" - } -}) + +@patch("api.modelmapper._model_map", {"provider1": {"model1": "mapped_model1", "model2": "mapped_model2"}}) class TestModelMapper(unittest.TestCase): def test_get_model_with_existing_model(self): result = get_model("provider1", "model1", "fallback_model") self.assertEqual(result, "mapped_model1") - @patch("api.modelmapper._model_map", { - "provider1": { - "model1": "mapped_model1", - "fallback_model": "fallback_model", - } - }) - + @patch( + "api.modelmapper._model_map", + { + "provider1": { + "model1": "mapped_model1", + "fallback_model": "fallback_model", + } + }, + ) def test_get_model_with_case_insensitivity(self): result = get_model("PROVIDER1", "MODEL1:latest", "fallback_model") self.assertEqual(result, "mapped_model1") @@ -34,12 +33,11 @@ def test_get_model_with_fallback(self): @patch("os.path.abspath", return_value="/mocked/path/modelmapper.py") def test_load_model_map(self, mock_abspath, mock_dirname, mock_join, mock_open_file): import api.modelmapper as modelmapper # <- directly access the module + modelmapper._model_map = None # Reset the actual global used by load_model_map modelmapper.load_model_map() - self.assertEqual( - modelmapper._model_map, - {"provider1": {"model1": "mapped_model1"}} - ) - + self.assertEqual(modelmapper._model_map, {"provider1": {"model1": "mapped_model1"}}) + + if __name__ == "__main__": unittest.main() diff --git a/src/requirements.txt b/src/requirements.txt index f609ce90..76da272b 100644 --- a/src/requirements.txt +++ b/src/requirements.txt @@ -17,3 +17,7 @@ google-auth-oauthlib==1.0.0 google-auth-httplib2==0.1.0 google-api-python-client==2.108.0 google-cloud-aiplatform==1.90.0 + +# linter +ruff>=0.12.5 +pre-commit>=2.20.0