Skip to content
Merged
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
2 changes: 1 addition & 1 deletion .github/workflows/lint.yml
Original file line number Diff line number Diff line change
Expand Up @@ -25,5 +25,5 @@ jobs:
with:
python-version: "3.12"
cache: pip
- run: pip install -r requirements.txt mypy
- run: pip install -r requirements.txt mypy==1.19.1
- run: mypy tee_gateway
7 changes: 7 additions & 0 deletions Makefile
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,12 @@ test-local:
# export OPENAI_API_KEY=... ANTHROPIC_API_KEY=... etc.
python3 -m tee_gateway

.PHONY: lint
lint:
python3 -m ruff format .
python3 -m ruff check .
python3 -m mypy tee_gateway

.PHONY: mypy
mypy:
python3 -m mypy tee_gateway
Expand All @@ -103,6 +109,7 @@ help:
@echo " make get-tls-cert - Print the nitriding TLS certificate"
@echo ""
@echo " make test-local - Run server locally without TEE (development)"
@echo " make lint - Run ruff check, ruff format --check, and mypy"
@echo " make mypy - Run mypy type checker on tee_gateway"
@echo ""
@echo " LLM endpoints (/v1/chat/completions, /v1/completions) require x402"
Expand Down
1 change: 1 addition & 0 deletions tests/dev-requirements.txt → dev-requirements.txt
Original file line number Diff line number Diff line change
@@ -1 +1,2 @@
mypy>=1.13.0
ruff>=0.9.1
3 changes: 1 addition & 2 deletions examples/verify_attestation.py
Original file line number Diff line number Diff line change
Expand Up @@ -74,7 +74,7 @@ def get_attestation(url: str, nonce: str) -> str:
return None

# Return the output of the curl command
if result.stdout == None:
if result.stdout is None:
print("Curl command result was None")

return result.stdout
Expand Down Expand Up @@ -197,7 +197,6 @@ def verify_attestation_doc(attestation_string: str) -> None:
cert_public_numbers = cert.get_pubkey().to_cryptography_key().public_numbers()
x = cert_public_numbers.x
y = cert_public_numbers.y
curve = cert_public_numbers.curve

x = long_to_bytes(x)
y = long_to_bytes(y)
Expand Down
10 changes: 3 additions & 7 deletions tee_gateway/__main__.py
Original file line number Diff line number Diff line change
Expand Up @@ -24,10 +24,11 @@
from x402v2.http.types import RouteConfig
from x402v2.mechanisms.evm.exact import ExactEvmServerScheme
from x402v2.mechanisms.evm.upto import UptoEvmServerScheme
from x402v2.schemas import AssetAmount, Network
from x402v2.schemas import AssetAmount
from x402v2.server import x402ResourceServerSync
from x402v2.session import SessionStore
import x402v2.http.middleware.flask as x402_flask

from .util import dynamic_session_cost_calculator
from .definitions import (
EVM_NETWORK,
Expand All @@ -38,6 +39,7 @@
CHAT_COMPLETIONS_USDC_AMOUNT,
CHAT_COMPLETIONS_OPG_AMOUNT,
COMPLETIONS_USDC_AMOUNT,
FACILITATOR_URL,
)

# Configure logging
Expand Down Expand Up @@ -82,12 +84,6 @@ def _shutdown_heartbeat():

atexit.register(_shutdown_heartbeat)


EVM_NETWORK: Network = "eip155:10740"
BASE_TESTNET_NETWORK: Network = "eip155:84532"
EVM_PAYMENT_ADDRESS = "0x40eFb45552EDfB2502D90A657a8ab41F03ec460d"
FACILITATOR_URL = os.getenv("FACILITATOR_URL", "https://facilitator.memchat.io")

facilitator = HTTPFacilitatorClientSync(FacilitatorConfig(url=FACILITATOR_URL))
server = x402ResourceServerSync(facilitator)
store = SessionStore()
Expand Down
7 changes: 6 additions & 1 deletion tee_gateway/definitions.py
Original file line number Diff line number Diff line change
Expand Up @@ -10,11 +10,16 @@

import os

# ---------------------------------------------------------------------------
# X402 Facilitator
# ---------------------------------------------------------------------------
FACILITATOR_URL = os.getenv("FACILITATOR_URL", "https://facilitator.memchat.io")

# ---------------------------------------------------------------------------
# Network IDs (EIP-155 chain identifiers)
# ---------------------------------------------------------------------------

# Lavanet EVM — where USDC payments are accepted
# OG EVM — where USDC payments are accepted
EVM_NETWORK: str = "eip155:10740"

# Base Testnet — where OPG payments are accepted
Expand Down
22 changes: 14 additions & 8 deletions tee_gateway/llm_backend.py
Original file line number Diff line number Diff line change
Expand Up @@ -14,8 +14,14 @@
from functools import lru_cache

import httpx

from langchain_core.messages import HumanMessage, SystemMessage, AIMessage, ToolMessage
from pydantic import SecretStr
from langchain_core.messages import (
HumanMessage,
SystemMessage,
AIMessage,
ToolMessage,
BaseMessage,
)
from langchain_google_genai import ChatGoogleGenerativeAI
from langchain_openai import ChatOpenAI
from langchain_anthropic import ChatAnthropic
Expand Down Expand Up @@ -145,13 +151,13 @@ def get_chat_model_cached(model: str, temperature: float, max_tokens: int):

return ChatOpenAI(
model=api_name,
api_key=api_key,
temperature=effective_temp,
max_tokens=max_tokens,
http_client=openai_http_client,
api_key=SecretStr(api_key),
streaming=True,
stream_usage=True,
)
) # type: ignore [call-arg]

elif provider == "anthropic":
api_key = os.getenv("ANTHROPIC_API_KEY")
Expand All @@ -160,13 +166,13 @@ def get_chat_model_cached(model: str, temperature: float, max_tokens: int):

return ChatAnthropic(
model=api_name,
api_key=api_key,
api_key=SecretStr(api_key),
temperature=effective_temp,
max_tokens=max_tokens,
timeout=ANTHROPIC_TIMEOUT,
streaming=True,
stream_usage=True,
)
) # type: ignore [call-arg]

elif provider == "x-ai":
api_key = os.getenv("XAI_API_KEY")
Expand All @@ -175,7 +181,7 @@ def get_chat_model_cached(model: str, temperature: float, max_tokens: int):

return ChatXAI(
model=api_name,
api_key=api_key,
api_key=SecretStr(api_key),
temperature=effective_temp,
max_tokens=max_tokens,
http_client=xai_http_client,
Expand All @@ -189,7 +195,7 @@ def get_chat_model_cached(model: str, temperature: float, max_tokens: int):

def convert_messages(messages: list) -> List[Any]:
"""Convert OpenAI-format message objects or dicts to LangChain message objects."""
langchain_messages = []
langchain_messages: List[BaseMessage] = []

for msg in messages:
# Support both OpenAPI model objects and plain dicts
Expand Down
10 changes: 5 additions & 5 deletions tee_gateway/typing_utils.py
Original file line number Diff line number Diff line change
Expand Up @@ -5,15 +5,15 @@

def is_generic(klass):
"""Determine whether klass is a generic class"""
return type(klass) == typing.GenericMeta
return type(klass) is typing.GenericMeta

def is_dict(klass):
"""Determine whether klass is a Dict"""
return klass.__extra__ == dict
return klass.__extra__ is dict

def is_list(klass):
"""Determine whether klass is a List"""
return klass.__extra__ == list
return klass.__extra__ is list

else:

Expand All @@ -23,8 +23,8 @@ def is_generic(klass):

def is_dict(klass):
"""Determine whether klass is a Dict"""
return klass.__origin__ == dict
return klass.__origin__ is dict

def is_list(klass):
"""Determine whether klass is a List"""
return klass.__origin__ == list
return klass.__origin__ is list
6 changes: 3 additions & 3 deletions tee_gateway/util.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,7 +23,7 @@ def _deserialize(data, klass):

if klass in (int, float, str, bool, bytearray):
return _deserialize_primitive(data, klass)
elif klass == object:
elif klass is object:
return _deserialize_object(data)
elif klass == datetime.date:
return deserialize_date(data)
Expand Down Expand Up @@ -76,7 +76,7 @@ def deserialize_date(string):
return None

try:
from dateutil.parser import parse
from dateutil.parser import parse # type: ignore[import-untyped]

return parse(string).date()
except ImportError:
Expand All @@ -97,7 +97,7 @@ def deserialize_datetime(string):
return None

try:
from dateutil.parser import parse
from dateutil.parser import parse # type: ignore[import-untyped]

return parse(string)
except ImportError:
Expand Down
30 changes: 0 additions & 30 deletions tests/mypy.ini

This file was deleted.

3 changes: 1 addition & 2 deletions tests/test_server.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,7 @@
from typing import List, Dict
from unittest.mock import patch, MagicMock
from dotenv import load_dotenv
from fastapi.testclient import TestClient

# Add src/ to path so we can import server
sys.path.insert(0, str(Path(__file__).resolve().parent.parent / "src"))
Expand All @@ -25,8 +26,6 @@
if missing_keys:
raise ValueError(f"Missing required API keys in .env: {', '.join(missing_keys)}")

from fastapi.testclient import TestClient

# Mock the TEE key manager registration before importing the app
with patch("urllib.request.urlopen"):
from server import app
Expand Down
Loading