Readability changes.

This commit is contained in:
Hemang
2025-06-05 11:58:12 +02:00
committed by Hemang Sarkar
parent cd6d6a50b0
commit 42a9c1cc30
20 changed files with 204 additions and 210 deletions
+2 -4
View File
@@ -7,8 +7,6 @@ import subprocess
import sys
import time
from typing import Optional
from gateway.mcp import stdio as mcp_stdio
from gateway.mcp.log import mcp_log
@@ -64,7 +62,7 @@ def ensure_network_exists(network_name: str = "invariant-explorer-web") -> bool:
return False
def setup_guardrails(guardrails_file_path: Optional[str] = None) -> bool:
def setup_guardrails(guardrails_file_path: str | None = None) -> bool:
"""Configure guardrails if specified."""
if not guardrails_file_path:
return True
@@ -105,7 +103,7 @@ def build():
return False
def up(guardrails_file_path: Optional[str] = None):
def up(guardrails_file_path: str | None = None):
"""Set up the local server for the Invariant Gateway."""
# Ensure network exists
if not ensure_network_exists():
+5 -6
View File
@@ -1,6 +1,5 @@
"""Common Authorization functions used in the gateway."""
from typing import Tuple, Optional
from fastapi import HTTPException, Request
INVARIANT_AUTHORIZATION_HEADER = "invariant-authorization"
@@ -10,7 +9,7 @@ API_KEYS_SEPARATOR = ";invariant-auth="
def extract_guardrail_service_authorization_from_headers(
request: Request,
) -> Tuple[Optional[str], Optional[str]]:
) -> tuple[str | None, str | None]:
"""
Extracts the optional Invariant-Guardrails-Authorization authorization header from the request.
@@ -22,10 +21,10 @@ def extract_guardrail_service_authorization_from_headers(
def extract_authorization_from_headers(
request: Request,
dataset_name: Optional[str] = None,
llm_provider_api_key_header: Optional[str] = None,
llm_provider_fallback_api_key_headers: Optional[list[str]] = None,
) -> Tuple[Optional[str], Optional[str]]:
dataset_name: str | None = None,
llm_provider_api_key_header: str | None = None,
llm_provider_fallback_api_key_headers: list[str] | None = None,
) -> tuple[str | None, str | None]:
"""
Extracts the Invariant authorization and LLM Provider API key from the request headers.
+5 -4
View File
@@ -3,14 +3,14 @@
import asyncio
import os
import threading
from typing import Optional
import fastapi
from httpx import HTTPStatusError
from gateway.common.guardrails import Guardrail, GuardrailAction, GuardrailRuleSet
def extract_policy_from_headers(request: Optional[fastapi.Request]) -> Optional[str]:
def extract_policy_from_headers(request: fastapi.Request | None) -> str | None:
"""
Extracts the guardrailing policy from the request headers if present.
@@ -78,7 +78,7 @@ class GatewayConfigManager:
_lock = threading.Lock()
@classmethod
def get_config(cls, request: fastapi.Request = None) -> GatewayConfig:
def get_config(cls) -> GatewayConfig:
"""Initializes and returns the gateway configuration using double-checked locking."""
local_config = cls._config_instance
@@ -94,7 +94,7 @@ class GatewayConfigManager:
async def extract_guardrails_from_header(
request: fastapi.Request,
) -> Optional[GuardrailRuleSet]:
) -> GuardrailRuleSet | None:
"""
Extracts Invariant-Guardrails from the request header if provided, and returns a corresponding
GuardrailRuleSet. If no guardrails are provided, returns None.
@@ -114,3 +114,4 @@ async def extract_guardrails_from_header(
blocking_guardrails=guardrails,
logging_guardrails=[],
)
return None
+2 -3
View File
@@ -2,7 +2,6 @@
from dataclasses import dataclass
from enum import Enum
from typing import List
class GuardrailAction(str, Enum):
"""Enum representing the action to be taken for guardrail rules."""
@@ -25,5 +24,5 @@ class Guardrail:
class GuardrailRuleSet:
"""Grouped guardrail rules separated by their action."""
blocking_guardrails: List[Guardrail]
logging_guardrails: List[Guardrail]
blocking_guardrails: list[Guardrail]
logging_guardrails: list[Guardrail]
+16 -16
View File
@@ -1,7 +1,7 @@
"""Common Request context data class."""
from dataclasses import dataclass, field
from typing import Any, Dict, Optional
from typing import Any
import fastapi
@@ -16,18 +16,18 @@ from gateway.common.guardrails import GuardrailRuleSet, Guardrail, GuardrailActi
class RequestContext:
"""Structured context for a request. Must be created via `RequestContext.create()`."""
request_json: Dict[str, Any]
dataset_name: Optional[str] = None
request_json: dict[str, Any]
dataset_name: str | None = None
# authorization to use for invariant service like explorer
invariant_authorization: Optional[str] = None
invariant_authorization: str | None = None
# authorization to use for invariant guardrailing specifically
guardrail_authorization: Optional[str] = None
guardrail_authorization: str | None = None
# the set of guardrails to enforce for this request
guardrails: Optional[GuardrailRuleSet] = None
config: Dict[str, Any] = None
guardrails: GuardrailRuleSet | None = None
config: dict[str, Any] | None = None
# extra parameters available as input.<key> during guardrail evaluation
guardrails_parameters: Optional[Dict[str, Any]] = None
guardrails_parameters: dict[str, Any] | None = None
_created_via_factory: bool = field(
default=False, init=True, repr=False, compare=False
@@ -42,13 +42,13 @@ class RequestContext:
@classmethod
def create(
cls,
request_json: Dict[str, Any],
dataset_name: Optional[str] = None,
invariant_authorization: Optional[str] = None,
guardrails: Optional[GuardrailRuleSet] = None,
config: Optional[GatewayConfig] = None,
request: fastapi.Request = None,
guardrails_parameters: Optional[Dict[str, Any]] = None,
request_json: dict[str, Any],
dataset_name: str | None = None,
invariant_authorization: str | None = None,
guardrails: GuardrailRuleSet | None = None,
config: GatewayConfig | None = None,
request: fastapi.Request | None = None,
guardrails_parameters: dict[str, Any] | None = None,
) -> "RequestContext":
"""Creates a new RequestContext instance, applying default guardrails if needed."""
@@ -103,7 +103,7 @@ class RequestContext:
guardrails_parameters=guardrails_parameters,
)
def get_guardrailing_authorization(self) -> Optional[str]:
def get_guardrailing_authorization(self) -> str | None:
"""
Returns the authorization to use for the guardrailing service.
+9 -9
View File
@@ -2,7 +2,7 @@
import os
import json
from typing import Any, Dict, List
from typing import Any
import httpx
from fastapi import HTTPException
@@ -15,8 +15,8 @@ from invariant_sdk.types.annotations import AnnotationCreate
def create_annotations_from_guardrails_errors(
guardrails_errors: List[dict],
) -> List[AnnotationCreate]:
guardrails_errors: list[dict],
) -> list[AnnotationCreate]:
"""Create Explorer annotations from the guardrails errors."""
annotations = []
@@ -67,7 +67,7 @@ def create_annotations_from_guardrails_errors(
return remove_duplicates(annotations)
def remove_duplicates(annotations: List[AnnotationCreate]) -> List[AnnotationCreate]:
def remove_duplicates(annotations: list[AnnotationCreate]) -> list[AnnotationCreate]:
"""
Remove duplicate annotations based on content, address, and extra_metadata.
@@ -98,18 +98,18 @@ def get_explorer_api_url() -> str:
async def push_trace(
messages: List[List[Dict[str, Any]]],
messages: list[list[dict[str, Any]]],
dataset_name: str,
invariant_authorization: str,
annotations: List[List[AnnotationCreate]] = None,
metadata: List[Dict[str, Any]] = None,
annotations: list[list[AnnotationCreate]] | None = None,
metadata: list[dict[str, Any]] | None = None,
) -> PushTracesResponse:
"""Pushes traces to the dataset on the Invariant Explorer.
If a dataset with the given name does not exist, it will be created.
Args:
messages (List[List[Dict[str, Any]]]): List of messages to push.
messages (listlistdict[str, Any]]]): List of messages to push.
dataset_name (str): Name of the dataset.
invariant_authorization (str): Value of the
invariant-authorization header.
@@ -134,7 +134,7 @@ async def push_trace(
)
try:
return await client.push_trace(request)
except Exception as e:
except Exception as e: # pylint: disable=broad-except
print(f"Failed to push trace: {e}")
return {"error": str(e)}
+7 -7
View File
@@ -3,7 +3,7 @@
import asyncio
import os
import time
from typing import Any, Dict, List
from typing import Any
from functools import wraps
from fastapi import HTTPException
@@ -339,22 +339,22 @@ class InstrumentedResponse(InstrumentedStreamingResponse):
async def check_guardrails(
messages: List[Dict[str, Any]],
guardrails: List[Guardrail],
messages: list[dict[str, Any]],
guardrails: list[Guardrail],
context: RequestContext,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""
Checks guardrails on the list of messages.
This calls the batch check API of the Guardrails service.
Args:
messages (List[Dict[str, Any]]): List of messages to verify the guardrails against.
guardrails (List[Guardrail]): The guardrails to check against.
messages (list[dict[str, Any]]): List of messages to verify the guardrails against.
guardrails (list[Guardrail]): The guardrails to check against.
invariant_authorization (str): Value of the
invariant-authorization header.
Returns:
Dict: Response containing guardrail check results.
dict: Response containing guardrail check results.
"""
async with httpx.AsyncClient() as client:
url = os.getenv("GUARDRAILS_API_URL", DEFAULT_API_URL).rstrip("/")
+102 -103
View File
@@ -7,7 +7,7 @@ import getpass
import os
import random
import socket
from typing import Any, Optional
from typing import Any
from invariant_sdk.async_client import AsyncClient
from invariant_sdk.types.append_messages import AppendMessagesRequest
@@ -32,6 +32,105 @@ def user_and_host() -> str:
return f"{username}@{hostname}"
class McpAttributes(BaseModel):
"""
A Pydantic model to represent MCP attributes.
This can be initialized using HTTP headers for SSE and Streamable transports.
This can also be initialized using CLI arguments for the Stdio transport.
"""
push_explorer: bool
explorer_dataset: str
invariant_api_key: str | None = None
verbose: bool | None = False
metadata: dict[str, Any] = Field(default_factory=dict)
@classmethod
def from_request_headers(cls, headers: Headers) -> "McpAttributes":
"""
Create an instance from FastAPI request headers.
Args:
headers: FastAPI Request headers
Returns:
McpAttributes: An instance with values extracted from headers
"""
# Extract and process header values
project_name = headers.get("INVARIANT-PROJECT-NAME")
push_explorer_header = headers.get("PUSH-INVARIANT-EXPLORER", "false").lower()
invariant_api_key = headers.get("INVARIANT-API-KEY")
# Determine explorer_dataset
if project_name:
explorer_dataset = project_name
else:
explorer_dataset = f"mcp-capture-{random.randint(1, 100)}"
# Determine push_explorer
push_explorer = push_explorer_header == "true"
# Create and return instance
return cls(
push_explorer=push_explorer,
explorer_dataset=explorer_dataset,
invariant_api_key=invariant_api_key,
)
@classmethod
def from_cli_args(cls, cli_args: list) -> "McpAttributes":
"""
Create an instance from command line arguments.
Args:
cli_args: List of command line arguments
Returns:
McpAttributes: An instance with values extracted from CLI arguments
"""
parser = argparse.ArgumentParser(description="MCP Gateway")
parser.add_argument(
"--project-name",
help="Name of the Project from Invariant Explorer where we want to push the MCP traces. The guardrails are pulled from this project.",
type=str,
default=f"mcp-capture-{random.randint(1, 100)}",
)
parser.add_argument(
"--push-explorer",
help="Enable pushing traces to Invariant Explorer",
action="store_true",
)
parser.add_argument(
"--verbose",
help="Enable verbose logging",
action="store_true",
)
parser.add_argument(
"--failure-response-format",
help="The response format to use to communicate guardrail failures to the client (error: JSON-RPC error response; potentially invisible to the agent, content: JSON-RPC content response, visible to the agent)",
type=str,
default="error",
)
config, extra_args = parser.parse_known_args(cli_args)
metadata: dict[str, Any] = {}
for arg in extra_args:
assert "=" in arg, f"Invalid extra metadata argument: {arg}"
key, value = arg.split("=")
assert key.startswith(
"--metadata-"
), f"Invalid extra metadata argument: {arg}, must start with --metadata-"
key = key[len("--metadata-") :]
metadata[key] = value
return cls(
push_explorer=config.push_explorer,
explorer_dataset=config.project_name,
verbose=config.verbose,
metadata=metadata,
)
class McpSession(BaseModel):
"""
@@ -40,9 +139,9 @@ class McpSession(BaseModel):
session_id: str
messages: list[dict[str, Any]] = Field(default_factory=list)
attributes: Optional["McpAttributes"] = None
attributes: McpAttributes | None = None
id_to_method_mapping: dict[int, str] = Field(default_factory=dict)
trace_id: Optional[str] = None
trace_id: str | None = None
last_trace_length: int = 0
annotations: list[dict[str, Any]] = Field(default_factory=list)
guardrails: GuardrailRuleSet = Field(
@@ -260,106 +359,6 @@ class McpSession(BaseModel):
return messages
class McpAttributes(BaseModel):
"""
A Pydantic model to represent MCP attributes.
This can be initialized using HTTP headers for SSE and Streamable transports.
This can also be initialized using CLI arguments for the Stdio transport.
"""
push_explorer: bool
explorer_dataset: str
invariant_api_key: Optional[str] = None
verbose: Optional[bool] = False
metadata: dict[str, Any] = Field(default_factory=dict)
@classmethod
def from_request_headers(cls, headers: Headers) -> "McpAttributes":
"""
Create an instance from FastAPI request headers.
Args:
headers: FastAPI Request headers
Returns:
McpAttributes: An instance with values extracted from headers
"""
# Extract and process header values
project_name = headers.get("INVARIANT-PROJECT-NAME")
push_explorer_header = headers.get("PUSH-INVARIANT-EXPLORER", "false").lower()
invariant_api_key = headers.get("INVARIANT-API-KEY")
# Determine explorer_dataset
if project_name:
explorer_dataset = project_name
else:
explorer_dataset = f"mcp-capture-{random.randint(1, 100)}"
# Determine push_explorer
push_explorer = push_explorer_header == "true"
# Create and return instance
return cls(
push_explorer=push_explorer,
explorer_dataset=explorer_dataset,
invariant_api_key=invariant_api_key,
)
@classmethod
def from_cli_args(cls, cli_args: list) -> "McpAttributes":
"""
Create an instance from command line arguments.
Args:
cli_args: List of command line arguments
Returns:
McpAttributes: An instance with values extracted from CLI arguments
"""
parser = argparse.ArgumentParser(description="MCP Gateway")
parser.add_argument(
"--project-name",
help="Name of the Project from Invariant Explorer where we want to push the MCP traces. The guardrails are pulled from this project.",
type=str,
default=f"mcp-capture-{random.randint(1, 100)}",
)
parser.add_argument(
"--push-explorer",
help="Enable pushing traces to Invariant Explorer",
action="store_true",
)
parser.add_argument(
"--verbose",
help="Enable verbose logging",
action="store_true",
)
parser.add_argument(
"--failure-response-format",
help="The response format to use to communicate guardrail failures to the client (error: JSON-RPC error response; potentially invisible to the agent, content: JSON-RPC content response, visible to the agent)",
type=str,
default="error",
)
config, extra_args = parser.parse_known_args(cli_args)
metadata: dict[str, Any] = {}
for arg in extra_args:
assert "=" in arg, f"Invalid extra metadata argument: {arg}"
key, value = arg.split("=")
assert key.startswith(
"--metadata-"
), f"Invalid extra metadata argument: {arg}, must start with --metadata-"
key = key[len("--metadata-") :]
metadata[key] = value
return cls(
push_explorer=config.push_explorer,
explorer_dataset=config.project_name,
verbose=config.verbose,
metadata=metadata,
)
class McpSessionsManager:
"""
A class to manage MCP sessions and their messages.
+12 -12
View File
@@ -8,7 +8,7 @@ import json
import re
import uuid
from abc import ABC, abstractmethod
from typing import Any, Tuple
from typing import Any
from fastapi import Request, HTTPException
from gateway.common.guardrails import GuardrailAction
@@ -43,12 +43,12 @@ class McpTransportBase(ABC):
async def process_outgoing_request(
self, session_id: str, request_data: dict[str, Any]
) -> Tuple[dict[str, Any], bool]:
) -> tuple[dict[str, Any], bool]:
"""
Template method for processing outgoing requests to MCP server.
Returns:
Tuple[processed_request_data, is_blocked]
tuple[processed_request_data, is_blocked]
"""
# Update session with request information
session = self.session_store.get_session(session_id)
@@ -65,12 +65,12 @@ class McpTransportBase(ABC):
async def process_incoming_response(
self, session_id: str, response_data: dict[str, Any]
) -> Tuple[dict[str, Any], bool]:
) -> tuple[dict[str, Any], bool]:
"""
Template method for processing incoming responses from MCP server.
Returns:
Tuple[processed_response, is_blocked]
tuple[processed_response, is_blocked]
"""
# Update session with server information
session = self.session_store.get_session(session_id)
@@ -99,7 +99,7 @@ class McpTransportBase(ABC):
async def _intercept_outgoing_request(
self, session_id: str, request_data: dict[str, Any]
) -> Tuple[dict[str, Any], bool]:
) -> tuple[dict[str, Any], bool]:
"""Common request interception logic for guardrails."""
method = request_data.get(MCP_METHOD)
@@ -209,7 +209,7 @@ class McpTransportBase(ABC):
@staticmethod
async def hook_tool_call(
session_id: str, session_store: McpSessionsManager, request_body: dict
) -> Tuple[dict, bool]:
) -> tuple[dict, bool]:
"""
Hook to process the request JSON before sending it to the MCP server.
@@ -219,7 +219,7 @@ class McpTransportBase(ABC):
request_body (dict): The request JSON to be processed.
Returns:
Tuple[dict, bool]: A tuple hook tool call response as a dict and a boolean
tuple[dict, bool]: A tuple hook tool call response as a dict and a boolean
indicating whether the request was blocked. If the request is blocked, the
dict will contain an error message else it will contain the original request.
"""
@@ -270,7 +270,7 @@ class McpTransportBase(ABC):
session_store: McpSessionsManager,
response_body: dict,
is_tools_list=False,
) -> Tuple[dict, bool]:
) -> tuple[dict, bool]:
"""
Hook to process the response JSON after receiving it from the MCP server.
@@ -280,7 +280,7 @@ class McpTransportBase(ABC):
response_body (dict): The response JSON to be processed.
is_tools_list (bool): Flag to indicate if the response is from a tools/list call.
Returns:
Tuple[dict, bool]: A tuple containing the processed response JSON
tuple[dict, bool]: A tuple containing the processed response JSON
and a boolean indicating whether the response was blocked. If the response
is blocked, the dict will contain an error message else it will contain the
original response.
@@ -351,7 +351,7 @@ class McpTransportBase(ABC):
@staticmethod
async def intercept_response(
session_id: str, session_store: McpSessionsManager, response_body: dict
) -> Tuple[dict, bool]:
) -> tuple[dict, bool]:
"""
Intercept the response and check for guardrails.
This function is used to intercept responses and check for guardrails.
@@ -365,7 +365,7 @@ class McpTransportBase(ABC):
response_body (dict): The response JSON to be processed.
Returns:
Tuple[dict, bool]: A tuple containing the processed response JSON
tuple[dict, bool]: A tuple containing the processed response JSON
and a boolean indicating whether the response was blocked.
"""
session = session_store.get_session(session_id)
+4 -4
View File
@@ -3,7 +3,7 @@
import asyncio
import json
import re
from typing import Any, AsyncGenerator, Optional, Tuple
from typing import Any, AsyncGenerator
import httpx
from httpx_sse import aconnect_sse, ServerSentEvent
@@ -85,8 +85,8 @@ class SseTransport(McpTransportBase):
**kwargs,
) -> str:
"""Initialize or get existing SSE session."""
session_id: Optional[str] = kwargs.get("session_id", None)
session_attributes: Optional[McpAttributes] = kwargs.get(
session_id: str | None = kwargs.get("session_id", None)
session_attributes: McpAttributes | None = kwargs.get(
"session_attributes", None
)
if session_id and self.session_store.session_exists(session_id):
@@ -303,7 +303,7 @@ class SseTransport(McpTransportBase):
async def _handle_endpoint_event(
self, sse: ServerSentEvent, sse_header_attributes: McpAttributes
) -> Tuple[bytes, str]:
) -> tuple[bytes, str]:
"""Handle endpoint event and initialize session if needed."""
match = re.search(r"session_id=([^&\s]+)", sse.data)
session_id = match.group(1) if match else None
+2 -3
View File
@@ -7,7 +7,6 @@ import platform
import select
import subprocess
import sys
from typing import Optional, Tuple
from gateway.mcp.constants import UTF_8
from gateway.mcp.log import mcp_log, MCP_LOG_FILE
@@ -210,7 +209,7 @@ class StdioTransport(McpTransportBase):
async def _wait_for_stdin_input(
self, loop: asyncio.AbstractEventLoop, stdin_fd: int
) -> Tuple[Optional[bytes], str]:
) -> tuple[bytes | None, str]:
"""Platform-specific implementation to wait for and read input from stdin."""
if platform.system() == "Windows":
await asyncio.sleep(0.01)
@@ -261,7 +260,7 @@ async def create_stdio_transport_and_execute(
)
def split_args(args: list[str] = None) -> tuple[list[str], list[str]]:
def split_args(args: list[str] | None = None) -> tuple[list[str], list[str]]:
"""
Splits CLI arguments into two parts:
1. Arguments intended for the MCP gateway (everything before `--exec`)
+4 -4
View File
@@ -1,7 +1,7 @@
"""Gateway service to forward requests to the MCP Streamable HTTP servers"""
import json
from typing import Any, Optional
from typing import Any
import httpx
from httpx_sse import aconnect_sse
@@ -90,8 +90,8 @@ class StreamableTransport(McpTransportBase):
**kwargs,
) -> str:
"""Initialize streamable HTTP session."""
session_id: Optional[str] = kwargs.get("session_id", None)
session_attributes: Optional[McpAttributes] = kwargs.get(
session_id: str | None = kwargs.get("session_id", None)
session_attributes: McpAttributes | None = kwargs.get(
"session_attributes", None
)
is_initialization_request: bool = kwargs.get("is_initialization_request", False)
@@ -240,7 +240,7 @@ class StreamableTransport(McpTransportBase):
async def _process_non_init_request(
self, session_id: str, request_body: dict[str, Any]
) -> Optional[Response]:
) -> Response | None:
"""Process non-initialization requests for guardrails."""
processed_request, is_blocked = await self.process_outgoing_request(
session_id, request_body
+6 -6
View File
@@ -69,7 +69,7 @@ def validate_headers(x_api_key: str = Header(None)):
)
async def anthropic_v1_messages_gateway(
request: Request,
dataset_name: str = None, # This is None if the client doesn't want to push to Explorer
dataset_name: str | None = None, # This is None if the client doesn't want to push to Explorer
config: GatewayConfig = Depends(GatewayConfigManager.get_config), # pylint: disable=unused-argument
header_guardrails: GuardrailRuleSet = Depends(extract_guardrails_from_header),
):
@@ -167,7 +167,7 @@ async def get_guardrails_check_result(
async def push_to_explorer(
context: RequestContext,
merged_response: dict[str, Any],
guardrails_execution_result: Optional[dict] = None,
guardrails_execution_result: dict | None = None,
) -> None:
"""Pushes the full trace to the Invariant Explorer"""
guardrails_execution_result = guardrails_execution_result or {}
@@ -215,9 +215,9 @@ class InstrumentedAnthropicResponse(InstrumentedResponse):
self.anthropic_request: httpx.Request = anthropic_request
# response data
self.response: Optional[httpx.Response] = None
self.response_string: Optional[str] = None
self.response_json: Optional[dict[str, Any]] = None
self.response: httpx.Response | None = None
self.response_string: str | None = None
self.response_json: dict[str, Any] | None = None
# guardrailing response (if any)
self.guardrails_execution_result = {}
@@ -553,7 +553,7 @@ class InstrumentedAnthropicStreamingResponse(InstrumentedStreamingResponse):
"""Process the buffer and extract complete SSE events.
Returns:
Tuple[List[str], str]: A tuple containing a list of
tuple[list[str], str]: A tuple containing a list of
complete events and the remaining buffer with incomplete events.
"""
# Split on double newlines which separate SSE events
+7 -7
View File
@@ -2,7 +2,7 @@
import asyncio
import json
from typing import Any, Literal, Optional
from typing import Any, Literal
import httpx
from fastapi import APIRouter, Depends, HTTPException, Query, Request, Response
@@ -49,7 +49,7 @@ async def gemini_generate_content_gateway(
api_version: str,
model: str,
endpoint: str,
dataset_name: str = None, # This is None if the client doesn't want to push to Explorer
dataset_name: str | None = None, # This is None if the client doesn't want to push to Explorer
alt: str = Query(
None, title="Response Format", description="Set to 'sse' for streaming"
),
@@ -147,7 +147,7 @@ class InstrumentedStreamingGeminiResponse(InstrumentedStreamingResponse):
}
# guardrailing execution result (if any)
self.guardrails_execution_result: Optional[dict[str, Any]] = None
self.guardrails_execution_result: dict[str, Any] | None = None
def make_refusal(
self,
@@ -415,7 +415,7 @@ async def get_guardrails_check_result(
async def push_to_explorer(
context: RequestContext,
response_json: dict[str, Any],
guardrails_execution_result: Optional[dict] = None,
guardrails_execution_result: dict | None = None,
) -> None:
"""Pushes the full trace to the Invariant Explorer"""
guardrails_execution_result = guardrails_execution_result or {}
@@ -464,11 +464,11 @@ class InstrumentedGeminiResponse(InstrumentedResponse):
self.gemini_request: httpx.Request = gemini_request
# response data
self.response: Optional[httpx.Response] = None
self.response_json: Optional[dict[str, Any]] = None
self.response: httpx.Response | None = None
self.response_json: dict[str, Any] | None = None
# guardrails execution result (if any)
self.guardrails_execution_result: Optional[dict[str, Any]] = None
self.guardrails_execution_result: dict[str, Any] | None = None
async def on_start(self):
"""
+10 -10
View File
@@ -2,7 +2,7 @@
import asyncio
import json
from typing import Any, Optional
from typing import Any
import httpx
from fastapi import APIRouter, Depends, Header, HTTPException, Request, Response
@@ -62,14 +62,14 @@ def make_cors_response(request: Request, allow_methods: str) -> Response:
@gateway.options("/{dataset_name}/openai/chat/completions")
@gateway.options("/openai/chat/completions")
async def openai_chat_completions_options(request: Request, dataset_name: str = None):
async def openai_chat_completions_options(request: Request):
"""Enables CORS for the OpenAI chat completions endpoint"""
return make_cors_response(request, allow_methods="POST")
@gateway.options("/{dataset_name}/openai/models")
@gateway.options("/openai/models")
async def openai_models_options(request: Request, dataset_name: str = None):
async def openai_models_options(request: Request):
"""Enables CORS for the OpenAI models endpoint"""
return make_cors_response(request, allow_methods="GET")
@@ -78,7 +78,7 @@ async def openai_models_options(request: Request, dataset_name: str = None):
@gateway.get("/openai/models")
async def openai_models_gateway(
request: Request,
dataset_name: str = None, # This is None if the client doesn't want to push to Explorer
dataset_name: str | None = None, # This is None if the client doesn't want to push to Explorer
):
"""Proxy request to OpenAI /models endpoint"""
headers = {
@@ -112,7 +112,7 @@ async def openai_models_gateway(
)
async def openai_chat_completions_gateway(
request: Request,
dataset_name: str = None, # This is None if the client doesn't want to push to Explorer
dataset_name: str | None = None, # This is None if the client doesn't want to push to Explorer
config: GatewayConfig = Depends(GatewayConfigManager.get_config), # pylint: disable=unused-argument
header_guardrails: GuardrailRuleSet = Depends(extract_guardrails_from_header),
) -> Response:
@@ -182,7 +182,7 @@ class InstrumentedOpenAIStreamResponse(InstrumentedStreamingResponse):
self.open_ai_request: httpx.Request = open_ai_request
# guardrailing output (if any)
self.guardrails_execution_result: Optional[dict] = None
self.guardrails_execution_result: dict | None = None
# merged_response will be updated with the data from the chunks in the stream
# At the end of the stream, this will be sent to the explorer
@@ -486,7 +486,7 @@ def create_metadata(
async def push_to_explorer(
context: RequestContext,
merged_response: dict[str, Any],
guardrails_execution_result: Optional[dict] = None,
guardrails_execution_result: dict | None = None,
) -> None:
"""Pushes the merged response to the Invariant Explorer"""
# Only push the trace to explorer if the message is an end turn message
@@ -572,11 +572,11 @@ class InstrumentedOpenAIResponse(InstrumentedResponse):
self.open_ai_request: httpx.Request = open_ai_request
# request outputs
self.response: Optional[httpx.Response] = None
self.response_json: Optional[dict[str, Any]] = None
self.response: httpx.Response | None = None
self.response_json: dict[str, Any] | None = None
# guardrailing output (if any)
self.guardrails_execution_result: Optional[dict] = None
self.guardrails_execution_result: dict | None = None
async def on_start(self):
"""
@@ -7,7 +7,6 @@ import sys
import time
import uuid
from pathlib import Path
from typing import Dict, List
# Add integration folder (parent) to sys.path
sys.path.append(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
@@ -49,7 +48,7 @@ class WeatherAgent:
},
}
def get_response(self, messages: List[Dict]) -> List[Dict]:
def get_response(self, messages: list[dict]) -> list[dict]:
"""
Get the response from the agent for a given user query for weather.
"""
@@ -83,7 +82,7 @@ class WeatherAgent:
else:
return response_list
def get_streaming_response(self, messages: List[Dict]) -> List[Dict]:
def get_streaming_response(self, messages: list[dict]) -> list[dict]:
"""Get streaming response from the agent for a given user query for weather."""
response_list = []
@@ -11,7 +11,7 @@ async def run(
gateway_url: str,
tool_name: str,
tool_args: dict[str, Any],
headers: dict[str, str] = None,
headers: dict[str, str] | None = None,
) -> types.CallToolResult | types.ListToolsResult:
"""
Run the MCP client with the given parameters.
@@ -3,7 +3,7 @@
import os
from datetime import timedelta
from typing import Any, Optional
from typing import Any
from mcp import ClientSession, StdioServerParameters, types
from mcp.client.stdio import stdio_client
@@ -14,7 +14,7 @@ def _get_server_params(
project_name: str,
server_script_path: str,
push_to_explorer: bool,
metadata_keys: Optional[dict[str, str]] = None,
metadata_keys: dict[str, str] | None = None,
) -> StdioServerParameters:
args = [
"--from",
@@ -59,7 +59,7 @@ async def run(
push_to_explorer: bool,
tool_name: str,
tool_args: dict[str, Any],
metadata_keys: Optional[dict[str, str]] = None,
metadata_keys: dict[str, str] | None = None,
) -> types.CallToolResult | types.ListToolsResult:
"""
Main function to setup the MCP client and server.
@@ -12,7 +12,7 @@ async def run(
gateway_url: str,
tool_name: str,
tool_args: dict[str, Any],
headers: dict[str, str] = None,
headers: dict[str, str] | None = None,
) -> types.CallToolResult | types.ListToolsResult:
"""
Run the MCP client with the given parameters.
+4 -4
View File
@@ -2,7 +2,7 @@
import os
import uuid
from typing import Any, Dict, Literal, Optional
from typing import Any, Literal
from httpx import Client
from openai import OpenAI
@@ -62,8 +62,8 @@ def get_gemini_client(
async def create_dataset(
explorer_api_url: str,
invariant_authorization: str,
dataset_name: Optional[str] = None,
) -> Dict[str, Any]:
dataset_name: str | None = None,
) -> dict[str, Any]:
"""Create a dataset in the Explorer API."""
client = Client(base_url=explorer_api_url)
response = client.post(
@@ -85,7 +85,7 @@ async def add_guardrail_to_dataset(
policy: str,
action: Literal["block", "log"],
invariant_authorization: str,
) -> Dict[str, Any]:
) -> dict[str, Any]:
"""Add a guardrail to a dataset."""
client = Client(base_url=explorer_api_url)
response = client.post(