diff --git a/.github/workflows/lint.yml b/.github/workflows/lint.yml new file mode 100644 index 0000000..768b6dc --- /dev/null +++ b/.github/workflows/lint.yml @@ -0,0 +1,27 @@ +name: Lint + +on: + pull_request: + +jobs: + ruff: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: astral-sh/ruff-action@v3 + with: + args: check tee_gateway + - uses: astral-sh/ruff-action@v3 + with: + args: format --check tee_gateway + + mypy: + runs-on: ubuntu-latest + steps: + - uses: actions/checkout@v4 + - uses: actions/setup-python@v5 + with: + python-version: "3.12" + cache: pip + - run: pip install -r requirements.txt mypy + - run: mypy tee_gateway diff --git a/Makefile b/Makefile index e1cb4c2..cc657bb 100644 --- a/Makefile +++ b/Makefile @@ -86,6 +86,10 @@ test-local: # export OPENAI_API_KEY=... ANTHROPIC_API_KEY=... etc. python3 -m tee_gateway +.PHONY: mypy +mypy: + python3 -m mypy tee_gateway + .PHONY: help help: @echo "Available targets:" @@ -99,6 +103,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 mypy - Run mypy type checker on tee_gateway" @echo "" @echo " LLM endpoints (/v1/chat/completions, /v1/completions) require x402" @echo " payment headers. Use an x402-compatible client to call them." diff --git a/examples/verify_attestation.py b/examples/verify_attestation.py index 518beb7..2625def 100644 --- a/examples/verify_attestation.py +++ b/examples/verify_attestation.py @@ -17,57 +17,56 @@ nonce = "0123456789abcdef0123456789abcdef01234567" logging.basicConfig( - filename='verification_logs.log', + filename="verification_logs.log", level=logging.DEBUG, - format='%(asctime)s - %(levelname)s - %(message)s', - datefmt='%Y-%m-%d %H:%M:%S' + format="%(asctime)s - %(levelname)s - %(message)s", + datefmt="%Y-%m-%d %H:%M:%S", ) PCR_tuple = namedtuple("PCRs", ["PCR0", "PCR1", "PCR2"]) -# This library is based on richardfan1126: nitro-enclave-python-demo + +# This library is based on richardfan1126: nitro-enclave-python-demo # (https://github.com/richardfan1126/nitro-enclave-python-demo/tree/master) def get_pcrs() -> PCR_tuple: """ Gets expected PCR values from enclave measurements JSON, returns a tuple of the PCR values. """ - with open(measurement_path, 'r') as file: + with open(measurement_path, "r") as file: json_measurement_data = file.read() try: measurement_data = json.loads(json_measurement_data) - PCRs = PCR_tuple(measurement_data["Measurements"]["PCR0"], - measurement_data["Measurements"]["PCR1"], - measurement_data["Measurements"]["PCR2"]) - logging.debug("Given PCR measurements:\n" - "PCR0 %s\n" - "PCR1 %s\n" - "PCR2 %s\n", - PCRs.PCR0, - PCRs.PCR1, - PCRs.PCR2) + PCRs = PCR_tuple( + measurement_data["Measurements"]["PCR0"], + measurement_data["Measurements"]["PCR1"], + measurement_data["Measurements"]["PCR2"], + ) + logging.debug( + "Given PCR measurements:\nPCR0 %s\nPCR1 %s\nPCR2 %s\n", + PCRs.PCR0, + PCRs.PCR1, + PCRs.PCR2, + ) except json.JSONDecodeError as e: raise ValueError("Error reading measurement file for PCRs: %s" % e) - + return PCRs + def get_root_cert_pem() -> str: - with open(root_cert_path, 'r') as file: + with open(root_cert_path, "r") as file: return file.read() + def get_attestation(url: str, nonce: str) -> str: # Construct curl command - curl_command = [ - "curl", - "-k", - "-G", - url, - "--data-urlencode", - f"nonce={nonce}" - ] + curl_command = ["curl", "-k", "-G", url, "--data-urlencode", f"nonce={nonce}"] # Run the curl command and capture the output - result = subprocess.run(curl_command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True) + result = subprocess.run( + curl_command, stdout=subprocess.PIPE, stderr=subprocess.PIPE, text=True + ) # Check if the command was successful if result.returncode != 0: @@ -80,6 +79,7 @@ def get_attestation(url: str, nonce: str) -> str: return result.stdout + def verify_attestation_doc(attestation_string: str) -> None: """ Verify the attestation document @@ -104,17 +104,19 @@ def verify_attestation_doc(attestation_string: str) -> None: logging.debug("Loaded an attestation document") # Expose timestamp - timestamp_ms = doc_obj['timestamp'] + timestamp_ms = doc_obj["timestamp"] timestamp_s = timestamp_ms / 1000 logging.info("Attestation document timestamp (ms): %d", timestamp_ms) - logging.info("Attestation document timestamp (UTC): %s", - __import__('datetime').datetime.utcfromtimestamp(timestamp_s).isoformat()) + logging.info( + "Attestation document timestamp (UTC): %s", + __import__("datetime").datetime.utcfromtimestamp(timestamp_s).isoformat(), + ) # Get PCRs from attestation document - document_pcrs_arr = doc_obj['pcrs'] + document_pcrs_arr = doc_obj["pcrs"] # Expose user data - user_data = doc_obj['user_data'] + user_data = doc_obj["user_data"] logging.debug("Enclave generated user data: %s", user_data) prefix_length = 2 # Taken from Nitriding documentation hash_length = hashlib.sha256().digest_size # 32 bytes for a SHA256 hash @@ -129,14 +131,18 @@ def verify_attestation_doc(attestation_string: str) -> None: app_key_end = app_key_start + hash_length app_key_hash = user_data[app_key_start:app_key_end] - logging.info("Enclave returned TLS key: %s", base64.b64encode(tls_key_hash).decode('utf-8')) - logging.info("Enclave returned app key: %s", base64.b64encode(app_key_hash).decode('utf-8')) + logging.info( + "Enclave returned TLS key: %s", base64.b64encode(tls_key_hash).decode("utf-8") + ) + logging.info( + "Enclave returned app key: %s", base64.b64encode(app_key_hash).decode("utf-8") + ) # Expose public key # TODO (kyle): Write API to expose public key to inference node # This will be needed for the sequencer to encrypt # input data. - public_key = doc_obj['public_key'] + public_key = doc_obj["public_key"] logging.debug("Enclave generated public key: %s", public_key) ## Validating Attestation document ## @@ -149,35 +155,42 @@ def verify_attestation_doc(attestation_string: str) -> None: # Get PCR hexcode doc_pcr = document_pcrs_arr[index].hex() - logging.debug("PCR%s:\n" - "Attestation value: %s\n" - "Expected PCR: %s", - index, - doc_pcr, - expected_pcr) + logging.debug( + "PCR%s:\nAttestation value: %s\nExpected PCR: %s", + index, + doc_pcr, + expected_pcr, + ) # Check if PCR match if expected_pcr != doc_pcr: - logging.warn("PCRs do not match:\n" - "Attestation PCR%s: %s\n" - "Expected PCR%s: %s", - index, doc_pcr, - index, expected_pcr) + logging.warn( + "PCRs do not match:\nAttestation PCR%s: %s\nExpected PCR%s: %s", + index, + doc_pcr, + index, + expected_pcr, + ) raise Exception("PCR%s does not match" % index) logging.debug("Validating nonce") # Check that nonce matches - attestation_nonce = doc_obj['nonce'].hex() + attestation_nonce = doc_obj["nonce"].hex() logging.info("Received nonce is %s", attestation_nonce) logging.info("Given nonce is %s", nonce) if attestation_nonce != nonce: - raise Exception(f"Attestation nonce: {attestation_nonce}, did not match given nonce: {nonce}") + raise Exception( + f"Attestation nonce: {attestation_nonce}, did not match given nonce: {nonce}" + ) - # 2. Validate Signature + # 2. Validate Signature logging.debug("Validating signature of attestation document") # Get signing certificate from attestation document - logging.debug("Getting signing certificate from attestation document:\n %s", doc_obj['certificate']) - cert = crypto.load_certificate(crypto.FILETYPE_ASN1, doc_obj['certificate']) + logging.debug( + "Getting signing certificate from attestation document:\n %s", + doc_obj["certificate"], + ) + cert = crypto.load_certificate(crypto.FILETYPE_ASN1, doc_obj["certificate"]) # Get the key parameters from the cert public key logging.debug("Creating EC2 key from the signing certificates public key") @@ -190,14 +203,14 @@ def verify_attestation_doc(attestation_string: str) -> None: y = long_to_bytes(y) # Create the EC2 key from public key parameters - key = EC2(alg = CoseAlgorithms.ES384, x = x, y = y, crv = CoseEllipticCurves.P_384) + key = EC2(alg=CoseAlgorithms.ES384, x=x, y=y, crv=CoseEllipticCurves.P_384) # Get the protected header from attestation document phdr = cbor2.loads(data[0]) # Construct the Sign1 message logging.debug("Constructing Sign1 message from the attestation document") - msg = cose.Sign1Message(phdr = phdr, uhdr = data[1], payload = doc) + msg = cose.Sign1Message(phdr=phdr, uhdr=data[1], payload=doc) msg.signature = data[3] # Verify the signature using the EC2 key @@ -207,8 +220,10 @@ def verify_attestation_doc(attestation_string: str) -> None: logging.debug("Signature of attestation document verified") # 3. Validate signing certificate PKI - logging.debug("Verifying the certificate of the attestation document " - "is signed by the root certificate of the AWS Nitro Attestation PKI") + logging.debug( + "Verifying the certificate of the attestation document " + "is signed by the root certificate of the AWS Nitro Attestation PKI" + ) if root_cert_pem is not None: # Create an X509Store object for the CA bundles store = crypto.X509Store() @@ -219,13 +234,13 @@ def verify_attestation_doc(attestation_string: str) -> None: # Get the CA bundle from attestation document and store into X509Store # Except the first certificate, which is the root certificate - for _cert_binary in doc_obj['cabundle'][1:]: + for _cert_binary in doc_obj["cabundle"][1:]: _cert = crypto.load_certificate(crypto.FILETYPE_ASN1, _cert_binary) store.add_cert(_cert) # Get the X509Store context store_ctx = crypto.X509StoreContext(store, cert) - + # Validate the certificate # If the cert is invalid, it will raise exception store_ctx.verify_certificate() @@ -234,10 +249,11 @@ def verify_attestation_doc(attestation_string: str) -> None: print("Verification successful") return + if __name__ == "__main__": print("Starting verification") attestation_str = get_attestation(enclave_url, nonce) print("Attestation string:\n", attestation_str) - verify_attestation_doc(attestation_str) \ No newline at end of file + verify_attestation_doc(attestation_str) diff --git a/examples/verify_signature_example.py b/examples/verify_signature_example.py index f248a69..29b8e08 100644 --- a/examples/verify_signature_example.py +++ b/examples/verify_signature_example.py @@ -22,17 +22,16 @@ -----END PUBLIC KEY-----""" public_key = serialization.load_pem_public_key( - public_key_pem.encode('utf-8'), - backend=default_backend() + public_key_pem.encode("utf-8"), backend=default_backend() ) # Parse system_fingerprint from your result system_fingerprint = "{'response_signature': 'PLyCgScL1Jr6OSb7wazEbor4yhBYJpauuqmsZJBoRNrpYl0sJ3ct472IminGRcfGGF1sBNB9YU6lKiWsJRnygJIufQ+yKt6a14QxrtjYp0F2LKCIvjIzveVnHs6oQQa9hz8VJFqSO/QLa4quw1GjYJHo+2fy8JPOPSBCXtmbHhBj4/7vSK53kQwJ0jld+LnpaAlURMxSaR49KOsbmAFCB9iR1pv292g0QOIY0hvlsNjH7HWaz+X1e2+Yytcl3eLP2IYIQqyVgJkm/U2zGb8ZZW10xxY+DbN+QlnHU9/SAq38n36zDbjLdZUhkgVTtht4vdn1wgFDSuEGQx4X5/nIHw==', 'request_signature': '3cd5e62557ea16dc77aef5c2c66188d180be259ac00f482de19896e78ebbf429', 'response_timestamp': '2025-12-16T08:20:26.491885+00:00'}" signatures = ast.literal_eval(system_fingerprint) -response_sig_b64 = signatures['response_signature'] -request_hash = signatures['request_signature'] -timestamp_iso = signatures['response_timestamp'] +response_sig_b64 = signatures["response_signature"] +request_hash = signatures["request_signature"] +timestamp_iso = signatures["response_timestamp"] response_signature = base64.b64decode(response_sig_b64) @@ -61,19 +60,19 @@ "content": "Hello!", "tool_calls": None, "tool_call_id": None, - "name": None + "name": None, } ], "max_tokens": 100, "temperature": 0.9, "stop": None, "tools": None, - "tool_choice": "auto" + "tool_choice": "auto", } # Compute request hash (matching server's compute_request_hash function) request_json = json.dumps(original_request, sort_keys=True) -computed_request_hash = hashlib.sha256(request_json.encode('utf-8')).hexdigest() +computed_request_hash = hashlib.sha256(request_json.encode("utf-8")).hexdigest() print(f"Request hash from response: {request_hash}") print(f"Computed from request: {computed_request_hash}") @@ -98,23 +97,20 @@ # Reconstruct the EXACT signed data structure (from server.py line 505-511) signed_data = { "finish_reason": finish_reason, - "message": { - "role": "assistant", - "content": message_content - }, + "message": {"role": "assistant", "content": message_content}, "model": model, "request_hash": request_hash, - "timestamp": timestamp_iso + "timestamp": timestamp_iso, } print("\nSigned data structure:") print(json.dumps(signed_data, indent=2, sort_keys=True)) # Create the signed message (with sort_keys=True to match server) -signed_message = json.dumps(signed_data, sort_keys=True).encode('utf-8') +signed_message = json.dumps(signed_data, sort_keys=True).encode("utf-8") print("\nSigned message (first 200 chars):") -print(signed_message[:200].decode('utf-8', errors='replace')) +print(signed_message[:200].decode("utf-8", errors="replace")) # Verify the signature using RSA-PSS try: @@ -122,12 +118,11 @@ response_signature, signed_message, padding.PSS( - mgf=padding.MGF1(hashes.SHA256()), - salt_length=padding.PSS.MAX_LENGTH + mgf=padding.MGF1(hashes.SHA256()), salt_length=padding.PSS.MAX_LENGTH ), - hashes.SHA256() + hashes.SHA256(), ) - + print("\n" + "=" * 70) print("✓✓✓ RESPONSE SIGNATURE VERIFIED! ✓✓✓") print("=" * 70) @@ -137,11 +132,11 @@ print(" 3. Not tampered with in transit") print("\nYou can cryptographically trust this message:") print(f' "{message_content}"') - + print("\n" + "=" * 70) print("✓✓✓ FULL TEE VERIFICATION SUCCESSFUL ✓✓✓") print("=" * 70) - + except Exception as e: print("\n✗ SIGNATURE VERIFICATION FAILED!") - print(f"Error: {e}") \ No newline at end of file + print(f"Error: {e}") diff --git a/mypy.ini b/mypy.ini new file mode 100644 index 0000000..abe3150 --- /dev/null +++ b/mypy.ini @@ -0,0 +1,3 @@ +[mypy] +python_version = 3.12 +ignore_missing_imports = True diff --git a/tee_gateway/__main__.py b/tee_gateway/__main__.py index 41f029f..99b87c4 100644 --- a/tee_gateway/__main__.py +++ b/tee_gateway/__main__.py @@ -106,7 +106,11 @@ def _shutdown_heartbeat(): price=AssetAmount( amount=CHAT_COMPLETIONS_USDC_AMOUNT, asset=USDC_ADDRESS, - extra={"name": "OUSDC", "version": "2", "assetTransferMethod": "permit2"}, + extra={ + "name": "OUSDC", + "version": "2", + "assetTransferMethod": "permit2", + }, ), network=EVM_NETWORK, ), @@ -116,7 +120,11 @@ def _shutdown_heartbeat(): price=AssetAmount( amount=CHAT_COMPLETIONS_OPG_AMOUNT, asset=BASE_OPG_ADDRESS, - extra={"name": "OPG", "version": "2", "assetTransferMethod": "permit2"}, + extra={ + "name": "OPG", + "version": "2", + "assetTransferMethod": "permit2", + }, ), network=BASE_TESTNET_NETWORK, ), @@ -159,29 +167,33 @@ def set_provider_keys(): with _keys_lock: if _keys_initialized: - return jsonify({"error": "Provider keys have already been initialized"}), 409 + return jsonify( + {"error": "Provider keys have already been initialized"} + ), 409 body = request.get_json(silent=True) if not body: return jsonify({"error": "JSON body required"}), 400 # Inject LLM provider keys into the environment - if body.get('openai_api_key'): - os.environ["OPENAI_API_KEY"] = body['openai_api_key'] - if body.get('google_api_key'): - os.environ["GOOGLE_API_KEY"] = body['google_api_key'] - if body.get('anthropic_api_key'): - os.environ["ANTHROPIC_API_KEY"] = body['anthropic_api_key'] - if body.get('xai_api_key'): - os.environ["XAI_API_KEY"] = body['xai_api_key'] + if body.get("openai_api_key"): + os.environ["OPENAI_API_KEY"] = body["openai_api_key"] + if body.get("google_api_key"): + os.environ["GOOGLE_API_KEY"] = body["google_api_key"] + if body.get("anthropic_api_key"): + os.environ["ANTHROPIC_API_KEY"] = body["anthropic_api_key"] + if body.get("xai_api_key"): + os.environ["XAI_API_KEY"] = body["xai_api_key"] # Inject heartbeat configuration into the environment - if body.get('heartbeat_contract_address'): - os.environ["HEARTBEAT_CONTRACT_ADDRESS"] = body['heartbeat_contract_address'] - if body.get('heartbeat_facilitator_url'): - os.environ["HEARTBEAT_FACILITATOR_URL"] = body['heartbeat_facilitator_url'] - if body.get('tee_heartbeat_interval'): - os.environ["TEE_HEARTBEAT_INTERVAL"] = str(body['tee_heartbeat_interval']) + if body.get("heartbeat_contract_address"): + os.environ["HEARTBEAT_CONTRACT_ADDRESS"] = body[ + "heartbeat_contract_address" + ] + if body.get("heartbeat_facilitator_url"): + os.environ["HEARTBEAT_FACILITATOR_URL"] = body["heartbeat_facilitator_url"] + if body.get("tee_heartbeat_interval"): + os.environ["TEE_HEARTBEAT_INTERVAL"] = str(body["tee_heartbeat_interval"]) def _key_status(env_var: str) -> str: return "set" if os.environ.get(env_var) else "NOT SET" @@ -189,12 +201,25 @@ def _key_status(env_var: str) -> str: logger.info("ENV check after injection:") logger.info(" OPENAI_API_KEY : %s", _key_status("OPENAI_API_KEY")) logger.info(" GOOGLE_API_KEY : %s", _key_status("GOOGLE_API_KEY")) - logger.info(" ANTHROPIC_API_KEY : %s", _key_status("ANTHROPIC_API_KEY")) + logger.info( + " ANTHROPIC_API_KEY : %s", _key_status("ANTHROPIC_API_KEY") + ) logger.info(" XAI_API_KEY : %s", _key_status("XAI_API_KEY")) - logger.info(" HEARTBEAT_CONTRACT_ADDRESS : %s", _key_status("HEARTBEAT_CONTRACT_ADDRESS")) - logger.info(" HEARTBEAT_FACILITATOR_URL : %s", _key_status("HEARTBEAT_FACILITATOR_URL")) - logger.info(" TEE_HEARTBEAT_INTERVAL : %s", os.environ.get("TEE_HEARTBEAT_INTERVAL", "900 (default)")) - logger.info(" HEARTBEAT_WALLET (TEE-gen) : %s", get_tee_keys().get_wallet_address()) + logger.info( + " HEARTBEAT_CONTRACT_ADDRESS : %s", + _key_status("HEARTBEAT_CONTRACT_ADDRESS"), + ) + logger.info( + " HEARTBEAT_FACILITATOR_URL : %s", + _key_status("HEARTBEAT_FACILITATOR_URL"), + ) + logger.info( + " TEE_HEARTBEAT_INTERVAL : %s", + os.environ.get("TEE_HEARTBEAT_INTERVAL", "900 (default)"), + ) + logger.info( + " HEARTBEAT_WALLET (TEE-gen) : %s", get_tee_keys().get_wallet_address() + ) # Rebuild HTTP clients with the new Authorization headers and clear # the model cache so subsequent requests use fresh instances. @@ -209,23 +234,30 @@ def _key_status(env_var: str) -> str: _keys_initialized = True providers_set = [ - p for p, k in { - "openai": body.get('openai_api_key'), - "google": body.get('google_api_key'), - "anthropic": body.get('anthropic_api_key'), - "xai": body.get('xai_api_key'), - }.items() if k + p + for p, k in { + "openai": body.get("openai_api_key"), + "google": body.get("google_api_key"), + "anthropic": body.get("anthropic_api_key"), + "xai": body.get("xai_api_key"), + }.items() + if k ] - heartbeat_configured = all([ - os.environ.get("HEARTBEAT_CONTRACT_ADDRESS"), - os.environ.get("HEARTBEAT_FACILITATOR_URL") or os.environ.get("FACILITATOR_URL"), - ]) + heartbeat_configured = all( + [ + os.environ.get("HEARTBEAT_CONTRACT_ADDRESS"), + os.environ.get("HEARTBEAT_FACILITATOR_URL") + or os.environ.get("FACILITATOR_URL"), + ] + ) logger.info("Provider API keys initialized for: %s", ", ".join(providers_set)) - return jsonify({ - "status": "ok", - "providers_initialized": providers_set, - "heartbeat_enabled": heartbeat_configured, - }), 200 + return jsonify( + { + "status": "ok", + "providers_initialized": providers_set, + "heartbeat_enabled": heartbeat_configured, + } + ), 200 def health(): @@ -279,14 +311,16 @@ def heartbeat_status(): def create_app(): app = connexion.App(__name__, specification_dir="./openapi/") app.app.json_encoder = encoder.JSONEncoder - app.add_api("openapi.yaml", - arguments={"title": "OpenAI API"}, - pythonic_params=True) + app.add_api("openapi.yaml", arguments={"title": "OpenAI API"}, pythonic_params=True) app.app.add_url_rule("/health", "health", health, methods=["GET"]) app.app.add_url_rule("/signing-key", "signing-key", signing_key, methods=["GET"]) - app.app.add_url_rule("/v1/keys", "set-provider-keys", set_provider_keys, methods=["POST"]) - app.app.add_url_rule("/heartbeat/status", "heartbeat-status", heartbeat_status, methods=["GET"]) + app.app.add_url_rule( + "/v1/keys", "set-provider-keys", set_provider_keys, methods=["POST"] + ) + app.app.add_url_rule( + "/heartbeat/status", "heartbeat-status", heartbeat_status, methods=["GET"] + ) # Initialize TEE here so it runs under both Gunicorn and direct execution. # This is the single TEEKeyManager instance — the same key both registers @@ -297,7 +331,6 @@ def create_app(): except Exception as e: logger.warning(f"TEE initialization failed (may not be in enclave): {e}") - return app.app @@ -311,6 +344,7 @@ def create_app(): # This patch ensures that non-payment 0-length requests can still bypass the middleware _original_read_body_bytes = x402_flask._read_body_bytes + def _patched_read_body_bytes(environ): try: content_length = int(environ.get("CONTENT_LENGTH") or 0) @@ -322,6 +356,7 @@ def _patched_read_body_bytes(environ): return _original_read_body_bytes(environ) + x402_flask._read_body_bytes = _patched_read_body_bytes payment_middleware( diff --git a/tee_gateway/controllers/chat_controller.py b/tee_gateway/controllers/chat_controller.py index 4e955e2..aed4e4e 100644 --- a/tee_gateway/controllers/chat_controller.py +++ b/tee_gateway/controllers/chat_controller.py @@ -5,9 +5,14 @@ import connexion from flask import Response +from typing import Any -from tee_gateway.models.create_chat_completion_request import CreateChatCompletionRequest -from tee_gateway.models.create_chat_completion_response import CreateChatCompletionResponse +from tee_gateway.models.create_chat_completion_request import ( + CreateChatCompletionRequest, +) +from tee_gateway.models.create_chat_completion_response import ( + CreateChatCompletionResponse, +) from tee_gateway.models import ( ChatCompletionRequestUserMessage, ChatCompletionRequestSystemMessage, @@ -31,11 +36,13 @@ def create_chat_completion(body): """Create a chat completion (streaming or non-streaming).""" if not connexion.request.is_json: return { - 'error': 'Unsupported Media Type', - 'message': 'Request must be application/json' + "error": "Unsupported Media Type", + "message": "Request must be application/json", }, 415 - chat_request: CreateChatCompletionRequest = _parse_chat_request(connexion.request.get_json()) + chat_request: CreateChatCompletionRequest = _parse_chat_request( + connexion.request.get_json() + ) if chat_request.stream: return _create_streaming_response(chat_request) @@ -52,11 +59,13 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): # Serialize request for hashing (canonical, deterministic) request_dict = _chat_request_to_dict(chat_request) - request_bytes = json.dumps(request_dict, sort_keys=True).encode('utf-8') + request_bytes = json.dumps(request_dict, sort_keys=True).encode("utf-8") model = get_chat_model_cached( model=chat_request.model, - temperature=float(chat_request.temperature) if chat_request.temperature is not None else 0.0, + temperature=float(chat_request.temperature) + if chat_request.temperature is not None + else 0.0, max_tokens=chat_request.max_tokens or 4096, ) @@ -65,8 +74,10 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): tools_list = [] for tool in chat_request.tools: if isinstance(tool, dict): - func = tool.get('function', {}) - tools_list.append({"type": tool.get('type', 'function'), "function": func}) + func = tool.get("function", {}) + tools_list.append( + {"type": tool.get("type", "function"), "function": func} + ) else: tools_list.append(tool) model = model.bind_tools(tools_list) @@ -76,14 +87,14 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): # Normalize content (Gemini may return a list of content parts) if isinstance(response.content, list): - content_str = ''.join( - item.get('text', '') if isinstance(item, dict) else str(item) + content_str = "".join( + item.get("text", "") if isinstance(item, dict) else str(item) for item in response.content ) else: content_str = response.content or "" - message_dict = { + message_dict: dict[str, Any] = { "role": "assistant", "content": content_str, } @@ -98,7 +109,7 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): "function": { "name": tc.get("name", ""), "arguments": json.dumps(tc.get("args", {})), - } + }, } for tc in response.tool_calls ] @@ -124,11 +135,13 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): "object": "chat.completion", "created": timestamp, "model": chat_request.model, - "choices": [{ - "index": 0, - "message": message_dict, - "finish_reason": finish_reason, - }], + "choices": [ + { + "index": 0, + "message": message_dict, + "finish_reason": finish_reason, + } + ], "tee_signature": signature, "tee_request_hash": input_hash_hex, "tee_output_hash": output_hash_hex, @@ -136,7 +149,9 @@ def _create_non_streaming_response(chat_request: CreateChatCompletionRequest): "tee_id": f"0x{tee_keys.get_tee_id()}", } - logger.debug(f"Response Final\n\tTEE Signature: {signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}") + logger.debug( + f"Response Final\n\tTEE Signature: {signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}" + ) if usage: openai_response["usage"] = usage @@ -159,11 +174,13 @@ def _create_streaming_response(chat_request: CreateChatCompletionRequest): buffer_tool_calls = provider in ["openai", "anthropic"] request_dict = _chat_request_to_dict(chat_request) - request_bytes = json.dumps(request_dict, sort_keys=True).encode('utf-8') + request_bytes = json.dumps(request_dict, sort_keys=True).encode("utf-8") model = get_chat_model_cached( model=chat_request.model, - temperature=float(chat_request.temperature) if chat_request.temperature is not None else 0.0, + temperature=float(chat_request.temperature) + if chat_request.temperature is not None + else 0.0, max_tokens=chat_request.max_tokens or 4096, ) @@ -171,8 +188,10 @@ def _create_streaming_response(chat_request: CreateChatCompletionRequest): tools_list = [] for tool in chat_request.tools: if isinstance(tool, dict): - func = tool.get('function', {}) - tools_list.append({"type": tool.get('type', 'function'), "function": func}) + func = tool.get("function", {}) + tools_list.append( + {"type": tool.get("type", "function"), "function": func} + ) else: tools_list.append(tool) model = model.bind_tools(tools_list) @@ -193,8 +212,10 @@ def generate(): if isinstance(chunk.content, str): content_str = chunk.content elif isinstance(chunk.content, list): - content_str = ''.join( - item.get('text', '') if isinstance(item, dict) else str(item) + content_str = "".join( + item.get("text", "") + if isinstance(item, dict) + else str(item) for item in chunk.content ) else: @@ -203,69 +224,104 @@ def generate(): if content_str: full_content += content_str data = { - "choices": [{ - "delta": {"content": content_str, "role": "assistant"}, - "index": 0, - "finish_reason": None - }], - "model": chat_request.model + "choices": [ + { + "delta": { + "content": content_str, + "role": "assistant", + }, + "index": 0, + "finish_reason": None, + } + ], + "model": chat_request.model, } yield f"data: {json.dumps(data)}\n\n" # --- Tool call chunks --- - if hasattr(chunk, 'tool_call_chunks') and chunk.tool_call_chunks: + if hasattr(chunk, "tool_call_chunks") and chunk.tool_call_chunks: finish_reason = "tool_calls" for tc_chunk in chunk.tool_call_chunks: - tc_index = tc_chunk.get('index', 0) + tc_index = tc_chunk.get("index", 0) if tc_index not in buffered_tool_calls: buffered_tool_calls[tc_index] = { - 'id': tc_chunk.get('id', ''), - 'type': 'function', - 'function': {'name': tc_chunk.get('name', ''), 'arguments': ''} + "id": tc_chunk.get("id", ""), + "type": "function", + "function": { + "name": tc_chunk.get("name", ""), + "arguments": "", + }, } - if tc_chunk.get('id'): - buffered_tool_calls[tc_index]['id'] = tc_chunk['id'] - if tc_chunk.get('name'): - buffered_tool_calls[tc_index]['function']['name'] = tc_chunk['name'] - if tc_chunk.get('args'): - args_value = tc_chunk['args'] + if tc_chunk.get("id"): + buffered_tool_calls[tc_index]["id"] = tc_chunk["id"] + if tc_chunk.get("name"): + buffered_tool_calls[tc_index]["function"]["name"] = ( + tc_chunk["name"] + ) + if tc_chunk.get("args"): + args_value = tc_chunk["args"] if isinstance(args_value, dict): args_str = json.dumps(args_value) elif isinstance(args_value, str): args_str = args_value else: args_str = str(args_value) - buffered_tool_calls[tc_index]['function']['arguments'] += args_str + buffered_tool_calls[tc_index]["function"][ + "arguments" + ] += args_str # For providers that don't need buffering, emit each fragment immediately if not buffer_tool_calls: - delta = {"role": "assistant", "tool_calls": [{"index": tc_index, "type": "function", "function": {}}]} - if tc_chunk.get('id'): - delta["tool_calls"][0]["id"] = tc_chunk['id'] - if tc_chunk.get('name'): - delta["tool_calls"][0]["function"]["name"] = tc_chunk['name'] - if tc_chunk.get('args'): - args_value = tc_chunk['args'] + delta = { + "role": "assistant", + "tool_calls": [ + { + "index": tc_index, + "type": "function", + "function": {}, + } + ], + } + if tc_chunk.get("id"): + delta["tool_calls"][0]["id"] = tc_chunk["id"] + if tc_chunk.get("name"): + delta["tool_calls"][0]["function"]["name"] = ( + tc_chunk["name"] + ) + if tc_chunk.get("args"): + args_value = tc_chunk["args"] if isinstance(args_value, dict): - delta["tool_calls"][0]["function"]["arguments"] = json.dumps(args_value) + delta["tool_calls"][0]["function"][ + "arguments" + ] = json.dumps(args_value) elif isinstance(args_value, str): - delta["tool_calls"][0]["function"]["arguments"] = args_value + delta["tool_calls"][0]["function"][ + "arguments" + ] = args_value else: - delta["tool_calls"][0]["function"]["arguments"] = str(args_value) + delta["tool_calls"][0]["function"][ + "arguments" + ] = str(args_value) if not delta["tool_calls"][0]["function"]: del delta["tool_calls"][0]["function"] data = { - "choices": [{"delta": delta, "index": 0, "finish_reason": None}], - "model": chat_request.model + "choices": [ + { + "delta": delta, + "index": 0, + "finish_reason": None, + } + ], + "model": chat_request.model, } yield f"data: {json.dumps(data)}\n\n" # --- Usage metadata --- - if hasattr(chunk, 'usage_metadata') and chunk.usage_metadata: + if hasattr(chunk, "usage_metadata") and chunk.usage_metadata: final_usage = chunk.usage_metadata # Flush buffered tool calls for OpenAI/Anthropic @@ -273,19 +329,23 @@ def generate(): for tc_index, tc in buffered_tool_calls.items(): delta = { "role": "assistant", - "tool_calls": [{ - "index": tc_index, - "id": tc['id'], - "type": "function", - "function": { - "name": tc['function']['name'], - "arguments": tc['function']['arguments'] + "tool_calls": [ + { + "index": tc_index, + "id": tc["id"], + "type": "function", + "function": { + "name": tc["function"]["name"], + "arguments": tc["function"]["arguments"], + }, } - }] + ], } data = { - "choices": [{"delta": delta, "index": 0, "finish_reason": None}], - "model": chat_request.model + "choices": [ + {"delta": delta, "index": 0, "finish_reason": None} + ], + "model": chat_request.model, } yield f"data: {json.dumps(data)}\n\n" @@ -294,7 +354,10 @@ def generate(): # signature covers the actual invocations. timestamp = int(time.time()) if finish_reason == "tool_calls" and buffered_tool_calls: - tool_calls_list = [buffered_tool_calls[k] for k in sorted(buffered_tool_calls.keys())] + tool_calls_list = [ + buffered_tool_calls[k] + for k in sorted(buffered_tool_calls.keys()) + ] output_content = json.dumps(tool_calls_list, sort_keys=True) else: output_content = full_content @@ -305,7 +368,9 @@ def generate(): tee_signature = tee_keys.sign_data(msg_hash) final_data = { - "choices": [{"delta": {}, "index": 0, "finish_reason": finish_reason}], + "choices": [ + {"delta": {}, "index": 0, "finish_reason": finish_reason} + ], "model": chat_request.model, "tee_signature": tee_signature, "tee_timestamp": timestamp, @@ -314,7 +379,9 @@ def generate(): "tee_id": f"0x{tee_keys.get_tee_id()}", } - logger.debug(f"Response Final\n\tTEE Signature: {tee_signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}") + logger.debug( + f"Response Final\n\tTEE Signature: {tee_signature}\n\tTEE request hash: {input_hash_hex}\n\tTEE output hash: {output_hash_hex}\n\tTEE timestamp: {timestamp}\n\tTEE ID: 0x{tee_keys.get_tee_id()}" + ) if final_usage: final_data["usage"] = { @@ -337,11 +404,11 @@ def generate(): return Response( generate(), - mimetype='text/event-stream', + mimetype="text/event-stream", headers={ - 'Cache-Control': 'no-cache', - 'Connection': 'keep-alive', - } + "Cache-Control": "no-cache", + "Connection": "keep-alive", + }, ) except Exception as e: @@ -353,6 +420,7 @@ def generate(): # Request parsing helpers # --------------------------------------------------------------------------- + def _chat_request_to_dict(chat_request: CreateChatCompletionRequest) -> dict: """Serialize a CreateChatCompletionRequest to a canonical dict for hashing.""" messages = [] @@ -360,67 +428,95 @@ def _chat_request_to_dict(chat_request: CreateChatCompletionRequest) -> dict: if isinstance(msg, ChatCompletionRequestSystemMessage): messages.append({"role": "system", "content": msg.content}) elif isinstance(msg, ChatCompletionRequestUserMessage): - messages.append({"role": "user", "content": msg.content if isinstance(msg.content, str) else str(msg.content)}) + messages.append( + { + "role": "user", + "content": msg.content + if isinstance(msg.content, str) + else str(msg.content), + } + ) elif isinstance(msg, ChatCompletionRequestAssistantMessage): m = {"role": "assistant", "content": msg.content or ""} if msg.tool_calls: m["tool_calls"] = [ - {"id": tc.id, "type": tc.type, "function": {"name": tc.function.name, "arguments": tc.function.arguments}} + { + "id": tc.id, + "type": tc.type, + "function": { + "name": tc.function.name, + "arguments": tc.function.arguments, + }, + } for tc in msg.tool_calls ] messages.append(m) elif isinstance(msg, ChatCompletionRequestToolMessage): - messages.append({"role": "tool", "content": msg.content, "tool_call_id": msg.tool_call_id}) + messages.append( + { + "role": "tool", + "content": msg.content, + "tool_call_id": msg.tool_call_id, + } + ) elif isinstance(msg, ChatCompletionRequestFunctionMessage): - messages.append({"role": "function", "content": msg.content, "name": msg.name}) + messages.append( + {"role": "function", "content": msg.content, "name": msg.name} + ) d = { "model": chat_request.model, "messages": messages, - "temperature": float(chat_request.temperature) if chat_request.temperature is not None else 0.0, + "temperature": float(chat_request.temperature) + if chat_request.temperature is not None + else 0.0, } if chat_request.max_tokens is not None: d["max_tokens"] = chat_request.max_tokens if chat_request.stop: d["stop"] = chat_request.stop if chat_request.tools: - d["tools"] = chat_request.tools if isinstance(chat_request.tools, list) else list(chat_request.tools) + d["tools"] = ( + chat_request.tools + if isinstance(chat_request.tools, list) + else list(chat_request.tools) + ) return d def _parse_chat_request(chat_request_dict: dict) -> CreateChatCompletionRequest: - messages = [_parse_message(msg) for msg in chat_request_dict.get('messages', [])] + messages = [_parse_message(msg) for msg in chat_request_dict.get("messages", [])] return CreateChatCompletionRequest( messages=messages, - model=chat_request_dict.get('model'), - frequency_penalty=chat_request_dict.get('frequency_penalty'), - logit_bias=chat_request_dict.get('logit_bias'), - max_tokens=chat_request_dict.get('max_tokens'), - n=chat_request_dict.get('n'), - presence_penalty=chat_request_dict.get('presence_penalty'), - response_format=chat_request_dict.get('response_format'), - seed=chat_request_dict.get('seed'), - stop=chat_request_dict.get('stop'), - stream=chat_request_dict.get('stream'), - temperature=chat_request_dict.get('temperature'), - top_p=chat_request_dict.get('top_p'), - tools=chat_request_dict.get('tools'), - tool_choice=chat_request_dict.get('tool_choice'), - user=chat_request_dict.get('user'), + model=chat_request_dict.get("model"), + frequency_penalty=chat_request_dict.get("frequency_penalty"), + logit_bias=chat_request_dict.get("logit_bias"), + max_tokens=chat_request_dict.get("max_tokens"), + n=chat_request_dict.get("n"), + presence_penalty=chat_request_dict.get("presence_penalty"), + response_format=chat_request_dict.get("response_format"), + seed=chat_request_dict.get("seed"), + stop=chat_request_dict.get("stop"), + stream=chat_request_dict.get("stream"), + temperature=chat_request_dict.get("temperature"), + top_p=chat_request_dict.get("top_p"), + tools=chat_request_dict.get("tools"), + tool_choice=chat_request_dict.get("tool_choice"), + user=chat_request_dict.get("user"), ) def _parse_message(message_dict: dict): - role = message_dict.get('role') - if role == 'user': + role = message_dict.get("role") + if role == "user": return ChatCompletionRequestUserMessage.from_dict(message_dict) - elif role == 'system': + elif role == "system": return ChatCompletionRequestSystemMessage.from_dict(message_dict) - elif role == 'assistant': + elif role == "assistant": return ChatCompletionRequestAssistantMessage.from_dict(message_dict) - elif role == 'tool': + elif role == "tool": return ChatCompletionRequestToolMessage.from_dict(message_dict) - elif role == 'function': + elif role == "function": return ChatCompletionRequestFunctionMessage.from_dict(message_dict) else: raise ValueError(f"Unknown role: {role}") diff --git a/tee_gateway/controllers/completions_controller.py b/tee_gateway/controllers/completions_controller.py index 4b4d42d..ba941fe 100644 --- a/tee_gateway/controllers/completions_controller.py +++ b/tee_gateway/controllers/completions_controller.py @@ -25,18 +25,22 @@ def create_completion(body): request_dict = { "model": body.model, "prompt": body.prompt, - "temperature": float(body.temperature) if body.temperature is not None else 0.0, + "temperature": float(body.temperature) + if body.temperature is not None + else 0.0, } if body.max_tokens is not None: request_dict["max_tokens"] = body.max_tokens if body.stop: request_dict["stop"] = body.stop - request_bytes = json.dumps(request_dict, sort_keys=True).encode('utf-8') + request_bytes = json.dumps(request_dict, sort_keys=True).encode("utf-8") model = get_chat_model_cached( model=body.model, - temperature=float(body.temperature) if body.temperature is not None else 0.0, + temperature=float(body.temperature) + if body.temperature is not None + else 0.0, max_tokens=body.max_tokens or 4096, ) @@ -58,11 +62,13 @@ def create_completion(body): "object": "text_completion", "created": timestamp, "model": body.model, - "choices": [{ - "text": response_content, - "index": 0, - "finish_reason": "stop", - }], + "choices": [ + { + "text": response_content, + "index": 0, + "finish_reason": "stop", + } + ], "usage": usage, "tee_signature": signature, "tee_request_hash": input_hash_hex, diff --git a/tee_gateway/controllers/defaults.py b/tee_gateway/controllers/defaults.py index 1a4f9df..e54a5e6 100644 --- a/tee_gateway/controllers/defaults.py +++ b/tee_gateway/controllers/defaults.py @@ -1,11 +1,9 @@ """ Default configuration for the OpenAPI server controllers. """ + import os # The internal LLM backend server (server.py running inside the enclave). # This is a temporary setup - controllers will eventually call LangChain directly. -HTTP_BACKEND_SERVER = os.getenv( - "LLM_BACKEND_URL", - "http://127.0.0.1:8001" -) \ No newline at end of file +HTTP_BACKEND_SERVER = os.getenv("LLM_BACKEND_URL", "http://127.0.0.1:8001") diff --git a/tee_gateway/controllers/security_controller.py b/tee_gateway/controllers/security_controller.py index 094d937..c2c199f 100644 --- a/tee_gateway/controllers/security_controller.py +++ b/tee_gateway/controllers/security_controller.py @@ -1,5 +1,3 @@ - - def info_from_ApiKeyAuth(token): """ OpenAPI security handler for ApiKeyAuth — intentional passthrough. @@ -15,4 +13,4 @@ def info_from_ApiKeyAuth(token): :return: Minimal token_info dict required by connexion :rtype: dict """ - return {'uid': 'user_id'} + return {"uid": "user_id"} diff --git a/tee_gateway/definitions.py b/tee_gateway/definitions.py index 2c9c1ae..0886abe 100644 --- a/tee_gateway/definitions.py +++ b/tee_gateway/definitions.py @@ -48,7 +48,7 @@ # Maps lowercase contract address → number of decimals for unit conversion. ASSET_DECIMALS_BY_ADDRESS: dict[str, int] = { - USDC_ADDRESS.lower(): 6, # USDC / OUSDC standard: 6 decimals + USDC_ADDRESS.lower(): 6, # USDC / OUSDC standard: 6 decimals BASE_OPG_ADDRESS.lower(): 18, # OPG: 18 decimals (ERC-20 standard) } diff --git a/tee_gateway/facilitator_api.py b/tee_gateway/facilitator_api.py index f550362..4eb8b60 100644 --- a/tee_gateway/facilitator_api.py +++ b/tee_gateway/facilitator_api.py @@ -1,7 +1,9 @@ """Pydantic models from TEE routing API to be used for facilitator.""" + from pydantic import BaseModel from typing import Optional, List, Dict, Any + class Message(BaseModel): role: str content: str @@ -56,7 +58,8 @@ class ChatResponse(BaseModel): class AttestationResponse(BaseModel): """TEE attestation document""" + public_key: str timestamp: str enclave_info: Dict[str, Any] - measurements: Optional[Dict] = None \ No newline at end of file + measurements: Optional[Dict] = None diff --git a/tee_gateway/heartbeat/heartbeat.py b/tee_gateway/heartbeat/heartbeat.py index 029fccb..a44fe7a 100644 --- a/tee_gateway/heartbeat/heartbeat.py +++ b/tee_gateway/heartbeat/heartbeat.py @@ -224,7 +224,7 @@ def create_heartbeat_service(tee_keys) -> Optional["HeartbeatService"]: "FACILITATOR_URL" ) - if not all([contract_address, facilitator_url]): + if not contract_address or not facilitator_url: logger.info( "Heartbeat disabled (set HEARTBEAT_CONTRACT_ADDRESS and " "HEARTBEAT_FACILITATOR_URL or FACILITATOR_URL to enable)" diff --git a/tee_gateway/llm_backend.py b/tee_gateway/llm_backend.py index eef0bb1..118e42e 100644 --- a/tee_gateway/llm_backend.py +++ b/tee_gateway/llm_backend.py @@ -118,7 +118,9 @@ def get_chat_model_cached(model: str, temperature: float, max_tokens: int): cfg = get_model_config(model) provider = cfg.provider api_name = cfg.api_name - effective_temp = cfg.force_temperature if cfg.force_temperature is not None else temperature + effective_temp = ( + cfg.force_temperature if cfg.force_temperature is not None else temperature + ) logger.info(f"Creating cached chat model - Provider: {provider}, Model: {api_name}") @@ -192,17 +194,17 @@ def convert_messages(messages: list) -> List[Any]: for msg in messages: # Support both OpenAPI model objects and plain dicts if isinstance(msg, dict): - role = msg.get('role', '').lower() - content = msg.get('content', '') or '' - tool_calls = msg.get('tool_calls') - tool_call_id = msg.get('tool_call_id') - name = msg.get('name') + role = msg.get("role", "").lower() + content = msg.get("content", "") or "" + tool_calls = msg.get("tool_calls") + tool_call_id = msg.get("tool_call_id") + name = msg.get("name") else: - role = getattr(msg, 'role', '').lower() - content = getattr(msg, 'content', '') or '' - tool_calls = getattr(msg, 'tool_calls', None) - tool_call_id = getattr(msg, 'tool_call_id', None) - name = getattr(msg, 'name', None) + role = getattr(msg, "role", "").lower() + content = getattr(msg, "content", "") or "" + tool_calls = getattr(msg, "tool_calls", None) + tool_call_id = getattr(msg, "tool_call_id", None) + name = getattr(msg, "name", None) if role == "system": langchain_messages.append(SystemMessage(content=content)) @@ -210,8 +212,8 @@ def convert_messages(messages: list) -> List[Any]: elif role == "user": # content may be a string or a list of content parts; handle both if isinstance(content, list): - content = ''.join( - part.get('text', '') if isinstance(part, dict) else str(part) + content = "".join( + part.get("text", "") if isinstance(part, dict) else str(part) for part in content ) langchain_messages.append(HumanMessage(content=content)) @@ -221,15 +223,15 @@ def convert_messages(messages: list) -> List[Any]: langchain_tool_calls = [] for tc in tool_calls: if isinstance(tc, dict): - func = tc.get('function', {}) - args = func.get('arguments', '{}') - tc_id = tc.get('id', '') - func_name = func.get('name', '') + func = tc.get("function", {}) + args = func.get("arguments", "{}") + tc_id = tc.get("id", "") + func_name = func.get("name", "") else: - func = getattr(tc, 'function', None) - args = func.arguments if func else '{}' - tc_id = getattr(tc, 'id', '') - func_name = func.name if func else '' + func = getattr(tc, "function", None) + args = func.arguments if func else "{}" + tc_id = getattr(tc, "id", "") + func_name = func.name if func else "" if isinstance(args, str): try: @@ -237,41 +239,49 @@ def convert_messages(messages: list) -> List[Any]: except json.JSONDecodeError: args = {} - langchain_tool_calls.append({ - "name": func_name, - "args": args, - "id": tc_id, - "type": "function", - }) - - langchain_messages.append(AIMessage( - content=content, - tool_calls=langchain_tool_calls, - )) + langchain_tool_calls.append( + { + "name": func_name, + "args": args, + "id": tc_id, + "type": "function", + } + ) + + langchain_messages.append( + AIMessage( + content=content, + tool_calls=langchain_tool_calls, + ) + ) else: langchain_messages.append(AIMessage(content=content)) elif role == "tool": - langchain_messages.append(ToolMessage( - content=content, - tool_call_id=tool_call_id or "", - name=name or "", - )) + langchain_messages.append( + ToolMessage( + content=content, + tool_call_id=tool_call_id or "", + name=name or "", + ) + ) elif role == "function": # Legacy function role: treat as tool message - langchain_messages.append(ToolMessage( - content=content, - tool_call_id="", - name=name or "", - )) + langchain_messages.append( + ToolMessage( + content=content, + tool_call_id="", + name=name or "", + ) + ) return langchain_messages def extract_usage(response) -> Optional[Dict[str, int]]: """Extract token usage from a LangChain response object.""" - if hasattr(response, 'usage_metadata') and response.usage_metadata: + if hasattr(response, "usage_metadata") and response.usage_metadata: meta = response.usage_metadata return { "prompt_tokens": meta.get("input_tokens", 0), diff --git a/tee_gateway/model_registry.py b/tee_gateway/model_registry.py index 1408797..f1bbe24 100644 --- a/tee_gateway/model_registry.py +++ b/tee_gateway/model_registry.py @@ -13,10 +13,10 @@ @dataclass(frozen=True) class ModelConfig: - provider: str # "openai" | "anthropic" | "google" | "x-ai" - api_name: str # model name sent to provider API - input_price_usd: Decimal # USD per token - output_price_usd: Decimal # USD per token + provider: str # "openai" | "anthropic" | "google" | "x-ai" + api_name: str # model name sent to provider API + input_price_usd: Decimal # USD per token + output_price_usd: Decimal # USD per token force_temperature: Optional[float] = None thinking_budget: Optional[int] = None @@ -190,39 +190,39 @@ class SupportedModel(Enum): # The "user-facing name" is what callers pass in the `model` field of requests. _MODEL_LOOKUP: dict[str, SupportedModel] = { # OpenAI - "gpt-4.1-2025-04-14": SupportedModel.GPT_4_1, - "gpt-4.1": SupportedModel.GPT_4_1, - "o4-mini": SupportedModel.O4_MINI, - "gpt-5": SupportedModel.GPT_5, - "gpt-5-mini": SupportedModel.GPT_5_MINI, - "gpt-5.2": SupportedModel.GPT_5_2, + "gpt-4.1-2025-04-14": SupportedModel.GPT_4_1, + "gpt-4.1": SupportedModel.GPT_4_1, + "o4-mini": SupportedModel.O4_MINI, + "gpt-5": SupportedModel.GPT_5, + "gpt-5-mini": SupportedModel.GPT_5_MINI, + "gpt-5.2": SupportedModel.GPT_5_2, # Anthropic - "claude-sonnet-4-5": SupportedModel.CLAUDE_SONNET_4_5, - "claude-sonnet-4-6": SupportedModel.CLAUDE_SONNET_4_6, - "claude-haiku-4-5": SupportedModel.CLAUDE_HAIKU_4_5, - "claude-opus-4-5": SupportedModel.CLAUDE_OPUS_4_5, - "claude-opus-4-6": SupportedModel.CLAUDE_OPUS_4_6, - "claude-3.7-sonnet": SupportedModel.CLAUDE_3_7_SONNET, - "claude-3.5-haiku": SupportedModel.CLAUDE_3_5_HAIKU, - "claude-4.0-sonnet": SupportedModel.CLAUDE_4_0_SONNET, + "claude-sonnet-4-5": SupportedModel.CLAUDE_SONNET_4_5, + "claude-sonnet-4-6": SupportedModel.CLAUDE_SONNET_4_6, + "claude-haiku-4-5": SupportedModel.CLAUDE_HAIKU_4_5, + "claude-opus-4-5": SupportedModel.CLAUDE_OPUS_4_5, + "claude-opus-4-6": SupportedModel.CLAUDE_OPUS_4_6, + "claude-3.7-sonnet": SupportedModel.CLAUDE_3_7_SONNET, + "claude-3.5-haiku": SupportedModel.CLAUDE_3_5_HAIKU, + "claude-4.0-sonnet": SupportedModel.CLAUDE_4_0_SONNET, # Google - "gemini-2.5-flash": SupportedModel.GEMINI_2_5_FLASH, - "gemini-2.5-pro": SupportedModel.GEMINI_2_5_PRO, - "gemini-2.5-flash-lite": SupportedModel.GEMINI_2_5_FLASH_LITE, - "gemini-3-pro-preview": SupportedModel.GEMINI_3_PRO_PREVIEW, - "gemini-3-flash-preview": SupportedModel.GEMINI_3_FLASH_PREVIEW, + "gemini-2.5-flash": SupportedModel.GEMINI_2_5_FLASH, + "gemini-2.5-pro": SupportedModel.GEMINI_2_5_PRO, + "gemini-2.5-flash-lite": SupportedModel.GEMINI_2_5_FLASH_LITE, + "gemini-3-pro-preview": SupportedModel.GEMINI_3_PRO_PREVIEW, + "gemini-3-flash-preview": SupportedModel.GEMINI_3_FLASH_PREVIEW, # xAI - "grok-4": SupportedModel.GROK_4, - "grok-4-fast": SupportedModel.GROK_4_FAST, - "grok-4-1-fast": SupportedModel.GROK_4_1_FAST, - "grok-4.1-fast": SupportedModel.GROK_4_1_FAST, - "grok-4-1-fast-non-reasoning": SupportedModel.GROK_4_1_FAST_NON_REASONING, - "grok-3-mini-beta": SupportedModel.GROK_3_MINI, - "grok-3-mini": SupportedModel.GROK_3_MINI, - "grok-3-beta": SupportedModel.GROK_3, - "grok-3": SupportedModel.GROK_3, - "grok-2-1212": SupportedModel.GROK_2, - "grok-2": SupportedModel.GROK_2, + "grok-4": SupportedModel.GROK_4, + "grok-4-fast": SupportedModel.GROK_4_FAST, + "grok-4-1-fast": SupportedModel.GROK_4_1_FAST, + "grok-4.1-fast": SupportedModel.GROK_4_1_FAST, + "grok-4-1-fast-non-reasoning": SupportedModel.GROK_4_1_FAST_NON_REASONING, + "grok-3-mini-beta": SupportedModel.GROK_3_MINI, + "grok-3-mini": SupportedModel.GROK_3_MINI, + "grok-3-beta": SupportedModel.GROK_3, + "grok-3": SupportedModel.GROK_3, + "grok-2-1212": SupportedModel.GROK_2, + "grok-2": SupportedModel.GROK_2, } # Build the rate card automatically from the enum (for backward compat with util.py) diff --git a/tee_gateway/models/__init__.py b/tee_gateway/models/__init__.py index c803370..ac0eceb 100644 --- a/tee_gateway/models/__init__.py +++ b/tee_gateway/models/__init__.py @@ -1,11 +1,25 @@ # flake8: noqa # import models into model package -from tee_gateway.models.chat_completion_request_assistant_message import ChatCompletionRequestAssistantMessage -from tee_gateway.models.chat_completion_request_function_message import ChatCompletionRequestFunctionMessage -from tee_gateway.models.chat_completion_request_system_message import ChatCompletionRequestSystemMessage -from tee_gateway.models.chat_completion_request_tool_message import ChatCompletionRequestToolMessage -from tee_gateway.models.chat_completion_request_user_message import ChatCompletionRequestUserMessage -from tee_gateway.models.create_chat_completion_request import CreateChatCompletionRequest -from tee_gateway.models.create_chat_completion_response import CreateChatCompletionResponse +from tee_gateway.models.chat_completion_request_assistant_message import ( + ChatCompletionRequestAssistantMessage, +) +from tee_gateway.models.chat_completion_request_function_message import ( + ChatCompletionRequestFunctionMessage, +) +from tee_gateway.models.chat_completion_request_system_message import ( + ChatCompletionRequestSystemMessage, +) +from tee_gateway.models.chat_completion_request_tool_message import ( + ChatCompletionRequestToolMessage, +) +from tee_gateway.models.chat_completion_request_user_message import ( + ChatCompletionRequestUserMessage, +) +from tee_gateway.models.create_chat_completion_request import ( + CreateChatCompletionRequest, +) +from tee_gateway.models.create_chat_completion_response import ( + CreateChatCompletionResponse, +) from tee_gateway.models.create_completion_request import CreateCompletionRequest from tee_gateway.models.create_completion_response import CreateCompletionResponse diff --git a/tee_gateway/models/chat_completion_request_assistant_message.py b/tee_gateway/models/chat_completion_request_assistant_message.py index 25ce696..c11b716 100644 --- a/tee_gateway/models/chat_completion_request_assistant_message.py +++ b/tee_gateway/models/chat_completion_request_assistant_message.py @@ -16,26 +16,34 @@ def __init__(self, id=None, type=None, function=None): class ChatCompletionRequestAssistantMessage(Model): - - def __init__(self, content=None, refusal=None, role=None, name=None, audio=None, tool_calls=None, function_call=None): # noqa: E501 + def __init__( + self, + content=None, + refusal=None, + role=None, + name=None, + audio=None, + tool_calls=None, + function_call=None, + ): # noqa: E501 self.openapi_types = { - 'content': object, - 'refusal': str, - 'role': str, - 'name': str, - 'audio': object, - 'tool_calls': object, - 'function_call': object, + "content": object, + "refusal": str, + "role": str, + "name": str, + "audio": object, + "tool_calls": object, + "function_call": object, } self.attribute_map = { - 'content': 'content', - 'refusal': 'refusal', - 'role': 'role', - 'name': 'name', - 'audio': 'audio', - 'tool_calls': 'tool_calls', - 'function_call': 'function_call', + "content": "content", + "refusal": "refusal", + "role": "role", + "name": "name", + "audio": "audio", + "tool_calls": "tool_calls", + "function_call": "function_call", } self._content = content @@ -47,29 +55,31 @@ def __init__(self, content=None, refusal=None, role=None, name=None, audio=None, self._function_call = function_call @classmethod - def from_dict(cls, dikt) -> 'ChatCompletionRequestAssistantMessage': - raw_tool_calls = dikt.get('tool_calls') + def from_dict(cls, dikt) -> "ChatCompletionRequestAssistantMessage": + raw_tool_calls = dikt.get("tool_calls") tool_calls = None if raw_tool_calls: tool_calls = [] for tc in raw_tool_calls: - func_raw = tc.get('function', {}) if isinstance(tc, dict) else {} - tool_calls.append(_ToolCall( - id=tc.get('id') if isinstance(tc, dict) else None, - type=tc.get('type') if isinstance(tc, dict) else None, - function=_ToolCallFunction( - name=func_raw.get('name'), - arguments=func_raw.get('arguments'), - ), - )) + func_raw = tc.get("function", {}) if isinstance(tc, dict) else {} + tool_calls.append( + _ToolCall( + id=tc.get("id") if isinstance(tc, dict) else None, + type=tc.get("type") if isinstance(tc, dict) else None, + function=_ToolCallFunction( + name=func_raw.get("name"), + arguments=func_raw.get("arguments"), + ), + ) + ) return cls( - content=dikt.get('content'), - refusal=dikt.get('refusal'), - role=dikt.get('role'), - name=dikt.get('name'), - audio=dikt.get('audio'), + content=dikt.get("content"), + refusal=dikt.get("refusal"), + role=dikt.get("role"), + name=dikt.get("name"), + audio=dikt.get("audio"), tool_calls=tool_calls, - function_call=dikt.get('function_call'), + function_call=dikt.get("function_call"), ) @property @@ -97,8 +107,9 @@ def role(self, role: str): allowed_values = ["assistant"] # noqa: E501 if role not in allowed_values: raise ValueError( - "Invalid value for `role` ({0}), must be one of {1}" - .format(role, allowed_values) + "Invalid value for `role` ({0}), must be one of {1}".format( + role, allowed_values + ) ) self._role = role diff --git a/tee_gateway/models/chat_completion_request_function_message.py b/tee_gateway/models/chat_completion_request_function_message.py index 34d6b4a..82660a8 100644 --- a/tee_gateway/models/chat_completion_request_function_message.py +++ b/tee_gateway/models/chat_completion_request_function_message.py @@ -22,24 +22,16 @@ def __init__(self, role=None, content=None, name=None): # noqa: E501 :param name: The name of this ChatCompletionRequestFunctionMessage. # noqa: E501 :type name: str """ - self.openapi_types = { - 'role': str, - 'content': str, - 'name': str - } - - self.attribute_map = { - 'role': 'role', - 'content': 'content', - 'name': 'name' - } + self.openapi_types = {"role": str, "content": str, "name": str} + + self.attribute_map = {"role": "role", "content": "content", "name": "name"} self._role = role self._content = content self._name = name @classmethod - def from_dict(cls, dikt) -> 'ChatCompletionRequestFunctionMessage': + def from_dict(cls, dikt) -> "ChatCompletionRequestFunctionMessage": """Returns the dict as a model :param dikt: A dict. @@ -72,8 +64,9 @@ def role(self, role: str): allowed_values = ["function"] # noqa: E501 if role not in allowed_values: raise ValueError( - "Invalid value for `role` ({0}), must be one of {1}" - .format(role, allowed_values) + "Invalid value for `role` ({0}), must be one of {1}".format( + role, allowed_values + ) ) self._role = role diff --git a/tee_gateway/models/chat_completion_request_system_message.py b/tee_gateway/models/chat_completion_request_system_message.py index 56ca278..30048f5 100644 --- a/tee_gateway/models/chat_completion_request_system_message.py +++ b/tee_gateway/models/chat_completion_request_system_message.py @@ -5,6 +5,7 @@ from tee_gateway.models.base_model import Model from tee_gateway import util + class ChatCompletionRequestSystemMessage(Model): """NOTE: This class is auto generated by OpenAPI Generator (https://openapi-generator.tech). @@ -21,24 +22,16 @@ def __init__(self, content=None, role=None, name=None): # noqa: E501 :param name: The name of this ChatCompletionRequestSystemMessage. # noqa: E501 :type name: str """ - self.openapi_types = { - 'content': object, - 'role': str, - 'name': str - } - - self.attribute_map = { - 'content': 'content', - 'role': 'role', - 'name': 'name' - } + self.openapi_types = {"content": object, "role": str, "name": str} + + self.attribute_map = {"content": "content", "role": "role", "name": "name"} self._content = content self._role = role self._name = name @classmethod - def from_dict(cls, dikt) -> 'ChatCompletionRequestSystemMessage': + def from_dict(cls, dikt) -> "ChatCompletionRequestSystemMessage": """Returns the dict as a model :param dikt: A dict. @@ -94,8 +87,9 @@ def role(self, role: str): allowed_values = ["system"] # noqa: E501 if role not in allowed_values: raise ValueError( - "Invalid value for `role` ({0}), must be one of {1}" - .format(role, allowed_values) + "Invalid value for `role` ({0}), must be one of {1}".format( + role, allowed_values + ) ) self._role = role diff --git a/tee_gateway/models/chat_completion_request_tool_message.py b/tee_gateway/models/chat_completion_request_tool_message.py index a1e30dd..dc8d1e1 100644 --- a/tee_gateway/models/chat_completion_request_tool_message.py +++ b/tee_gateway/models/chat_completion_request_tool_message.py @@ -5,6 +5,7 @@ from tee_gateway.models.base_model import Model from tee_gateway import util + class ChatCompletionRequestToolMessage(Model): """NOTE: This class is auto generated by OpenAPI Generator (https://openapi-generator.tech). @@ -21,16 +22,12 @@ def __init__(self, role=None, content=None, tool_call_id=None): # noqa: E501 :param tool_call_id: The tool_call_id of this ChatCompletionRequestToolMessage. # noqa: E501 :type tool_call_id: str """ - self.openapi_types = { - 'role': str, - 'content': object, - 'tool_call_id': str - } + self.openapi_types = {"role": str, "content": object, "tool_call_id": str} self.attribute_map = { - 'role': 'role', - 'content': 'content', - 'tool_call_id': 'tool_call_id' + "role": "role", + "content": "content", + "tool_call_id": "tool_call_id", } self._role = role @@ -38,7 +35,7 @@ def __init__(self, role=None, content=None, tool_call_id=None): # noqa: E501 self._tool_call_id = tool_call_id @classmethod - def from_dict(cls, dikt) -> 'ChatCompletionRequestToolMessage': + def from_dict(cls, dikt) -> "ChatCompletionRequestToolMessage": """Returns the dict as a model :param dikt: A dict. @@ -71,8 +68,9 @@ def role(self, role: str): allowed_values = ["tool"] # noqa: E501 if role not in allowed_values: raise ValueError( - "Invalid value for `role` ({0}), must be one of {1}" - .format(role, allowed_values) + "Invalid value for `role` ({0}), must be one of {1}".format( + role, allowed_values + ) ) self._role = role diff --git a/tee_gateway/models/chat_completion_request_user_message.py b/tee_gateway/models/chat_completion_request_user_message.py index 70d3cd9..24e106e 100644 --- a/tee_gateway/models/chat_completion_request_user_message.py +++ b/tee_gateway/models/chat_completion_request_user_message.py @@ -5,6 +5,7 @@ from tee_gateway.models.base_model import Model from tee_gateway import util + class ChatCompletionRequestUserMessage(Model): """NOTE: This class is auto generated by OpenAPI Generator (https://openapi-generator.tech). @@ -21,24 +22,16 @@ def __init__(self, content=None, role=None, name=None): # noqa: E501 :param name: The name of this ChatCompletionRequestUserMessage. # noqa: E501 :type name: str """ - self.openapi_types = { - 'content': object, - 'role': str, - 'name': str - } - - self.attribute_map = { - 'content': 'content', - 'role': 'role', - 'name': 'name' - } + self.openapi_types = {"content": object, "role": str, "name": str} + + self.attribute_map = {"content": "content", "role": "role", "name": "name"} self._content = content self._role = role self._name = name @classmethod - def from_dict(cls, dikt) -> 'ChatCompletionRequestUserMessage': + def from_dict(cls, dikt) -> "ChatCompletionRequestUserMessage": """Returns the dict as a model :param dikt: A dict. @@ -94,8 +87,9 @@ def role(self, role: str): allowed_values = ["user"] # noqa: E501 if role not in allowed_values: raise ValueError( - "Invalid value for `role` ({0}), must be one of {1}" - .format(role, allowed_values) + "Invalid value for `role` ({0}), must be one of {1}".format( + role, allowed_values + ) ) self._role = role diff --git a/tee_gateway/models/create_chat_completion_request.py b/tee_gateway/models/create_chat_completion_request.py index 09f21c6..a32a3d5 100644 --- a/tee_gateway/models/create_chat_completion_request.py +++ b/tee_gateway/models/create_chat_completion_request.py @@ -1,14 +1,39 @@ class CreateChatCompletionRequest: """Minimal data holder for chat completion requests.""" - def __init__(self, messages=None, model=None, store=False, reasoning_effort='medium', - metadata=None, frequency_penalty=0, logit_bias=None, logprobs=False, - top_logprobs=None, max_tokens=None, max_completion_tokens=None, n=1, - modalities=None, prediction=None, audio=None, presence_penalty=0, - response_format=None, seed=None, service_tier='auto', stop=None, - stream=False, stream_options=None, temperature=1, top_p=1, - tools=None, tool_choice=None, parallel_tool_calls=True, - user=None, function_call=None, functions=None): + def __init__( + self, + messages=None, + model=None, + store=False, + reasoning_effort="medium", + metadata=None, + frequency_penalty=0, + logit_bias=None, + logprobs=False, + top_logprobs=None, + max_tokens=None, + max_completion_tokens=None, + n=1, + modalities=None, + prediction=None, + audio=None, + presence_penalty=0, + response_format=None, + seed=None, + service_tier="auto", + stop=None, + stream=False, + stream_options=None, + temperature=1, + top_p=1, + tools=None, + tool_choice=None, + parallel_tool_calls=True, + user=None, + function_call=None, + functions=None, + ): self.messages = messages self.model = model self.store = store @@ -41,15 +66,39 @@ def __init__(self, messages=None, model=None, store=False, reasoning_effort='med self.functions = functions @classmethod - def from_dict(cls, dikt) -> 'CreateChatCompletionRequest': + def from_dict(cls, dikt) -> "CreateChatCompletionRequest": if not isinstance(dikt, dict): return dikt known = { - 'messages', 'model', 'store', 'reasoning_effort', 'metadata', - 'frequency_penalty', 'logit_bias', 'logprobs', 'top_logprobs', - 'max_tokens', 'max_completion_tokens', 'n', 'modalities', 'prediction', - 'audio', 'presence_penalty', 'response_format', 'seed', 'service_tier', - 'stop', 'stream', 'stream_options', 'temperature', 'top_p', 'tools', - 'tool_choice', 'parallel_tool_calls', 'user', 'function_call', 'functions', + "messages", + "model", + "store", + "reasoning_effort", + "metadata", + "frequency_penalty", + "logit_bias", + "logprobs", + "top_logprobs", + "max_tokens", + "max_completion_tokens", + "n", + "modalities", + "prediction", + "audio", + "presence_penalty", + "response_format", + "seed", + "service_tier", + "stop", + "stream", + "stream_options", + "temperature", + "top_p", + "tools", + "tool_choice", + "parallel_tool_calls", + "user", + "function_call", + "functions", } return cls(**{k: v for k, v in dikt.items() if k in known}) diff --git a/tee_gateway/models/create_chat_completion_response.py b/tee_gateway/models/create_chat_completion_response.py index 04deb3a..8d5c276 100644 --- a/tee_gateway/models/create_chat_completion_response.py +++ b/tee_gateway/models/create_chat_completion_response.py @@ -1,8 +1,17 @@ class CreateChatCompletionResponse: """Minimal data holder for chat completion responses (used for validation).""" - def __init__(self, id=None, choices=None, created=None, model=None, - service_tier=None, system_fingerprint=None, object=None, usage=None): + def __init__( + self, + id=None, + choices=None, + created=None, + model=None, + service_tier=None, + system_fingerprint=None, + object=None, + usage=None, + ): self.id = id self.choices = choices self.created = created @@ -13,16 +22,16 @@ def __init__(self, id=None, choices=None, created=None, model=None, self.usage = usage @classmethod - def from_dict(cls, dikt) -> 'CreateChatCompletionResponse': + def from_dict(cls, dikt) -> "CreateChatCompletionResponse": if not isinstance(dikt, dict): return dikt return cls( - id=dikt.get('id'), - choices=dikt.get('choices'), - created=dikt.get('created'), - model=dikt.get('model'), - service_tier=dikt.get('service_tier'), - system_fingerprint=dikt.get('system_fingerprint'), - object=dikt.get('object'), - usage=dikt.get('usage'), + id=dikt.get("id"), + choices=dikt.get("choices"), + created=dikt.get("created"), + model=dikt.get("model"), + service_tier=dikt.get("service_tier"), + system_fingerprint=dikt.get("system_fingerprint"), + object=dikt.get("object"), + usage=dikt.get("usage"), ) diff --git a/tee_gateway/models/create_completion_request.py b/tee_gateway/models/create_completion_request.py index c85b019..d073666 100644 --- a/tee_gateway/models/create_completion_request.py +++ b/tee_gateway/models/create_completion_request.py @@ -1,10 +1,27 @@ class CreateCompletionRequest: """Stub data holder for completion requests (endpoint not implemented).""" - def __init__(self, model=None, prompt=None, best_of=1, echo=False, - frequency_penalty=0, logit_bias=None, logprobs=None, max_tokens=16, - n=1, presence_penalty=0, seed=None, stop=None, stream=False, - stream_options=None, suffix=None, temperature=1, top_p=1, user=None): + def __init__( + self, + model=None, + prompt=None, + best_of=1, + echo=False, + frequency_penalty=0, + logit_bias=None, + logprobs=None, + max_tokens=16, + n=1, + presence_penalty=0, + seed=None, + stop=None, + stream=False, + stream_options=None, + suffix=None, + temperature=1, + top_p=1, + user=None, + ): self.model = model self.prompt = prompt self.best_of = best_of @@ -25,12 +42,27 @@ def __init__(self, model=None, prompt=None, best_of=1, echo=False, self.user = user @classmethod - def from_dict(cls, dikt) -> 'CreateCompletionRequest': + def from_dict(cls, dikt) -> "CreateCompletionRequest": if not isinstance(dikt, dict): return dikt known = { - 'model', 'prompt', 'best_of', 'echo', 'frequency_penalty', 'logit_bias', - 'logprobs', 'max_tokens', 'n', 'presence_penalty', 'seed', 'stop', - 'stream', 'stream_options', 'suffix', 'temperature', 'top_p', 'user', + "model", + "prompt", + "best_of", + "echo", + "frequency_penalty", + "logit_bias", + "logprobs", + "max_tokens", + "n", + "presence_penalty", + "seed", + "stop", + "stream", + "stream_options", + "suffix", + "temperature", + "top_p", + "user", } return cls(**{k: v for k, v in dikt.items() if k in known}) diff --git a/tee_gateway/models/create_completion_response.py b/tee_gateway/models/create_completion_response.py index b991dde..44f7cb8 100644 --- a/tee_gateway/models/create_completion_response.py +++ b/tee_gateway/models/create_completion_response.py @@ -1,8 +1,16 @@ class CreateCompletionResponse: """Stub data holder for completion responses (endpoint not implemented).""" - def __init__(self, id=None, choices=None, created=None, model=None, - system_fingerprint=None, object=None, usage=None): + def __init__( + self, + id=None, + choices=None, + created=None, + model=None, + system_fingerprint=None, + object=None, + usage=None, + ): self.id = id self.choices = choices self.created = created @@ -12,15 +20,15 @@ def __init__(self, id=None, choices=None, created=None, model=None, self.usage = usage @classmethod - def from_dict(cls, dikt) -> 'CreateCompletionResponse': + def from_dict(cls, dikt) -> "CreateCompletionResponse": if not isinstance(dikt, dict): return dikt return cls( - id=dikt.get('id'), - choices=dikt.get('choices'), - created=dikt.get('created'), - model=dikt.get('model'), - system_fingerprint=dikt.get('system_fingerprint'), - object=dikt.get('object'), - usage=dikt.get('usage'), + id=dikt.get("id"), + choices=dikt.get("choices"), + created=dikt.get("created"), + model=dikt.get("model"), + system_fingerprint=dikt.get("system_fingerprint"), + object=dikt.get("object"), + usage=dikt.get("usage"), ) diff --git a/tee_gateway/tee_manager.py b/tee_gateway/tee_manager.py index de97544..3d98b12 100644 --- a/tee_gateway/tee_manager.py +++ b/tee_gateway/tee_manager.py @@ -44,16 +44,14 @@ def _generate_keys(self): """Generate RSA key pair and derive the tee_id.""" logger.info("Generating TEE RSA key pair...") self.private_key = rsa.generate_private_key( - public_exponent=65537, - key_size=2048, - backend=default_backend() + public_exponent=65537, key_size=2048, backend=default_backend() ) self.public_key = self.private_key.public_key() self.public_key_pem = self.public_key.public_bytes( encoding=serialization.Encoding.PEM, - format=serialization.PublicFormat.SubjectPublicKeyInfo - ).decode('utf-8') + format=serialization.PublicFormat.SubjectPublicKeyInfo, + ).decode("utf-8") # tee_id = keccak256(abi.encodePacked(signingKey)) # signingKey is the DER-encoded public key (canonical binary form, no @@ -62,7 +60,7 @@ def _generate_keys(self): # the body, and keccak256-hash the resulting bytes. public_key_der = self.public_key.public_bytes( encoding=serialization.Encoding.DER, - format=serialization.PublicFormat.SubjectPublicKeyInfo + format=serialization.PublicFormat.SubjectPublicKeyInfo, ) self.tee_id = keccak(public_key_der).hex() @@ -81,11 +79,11 @@ def register_with_nitriding(self): try: public_key_der = self.public_key.public_bytes( encoding=serialization.Encoding.DER, - format=serialization.PublicFormat.SubjectPublicKeyInfo + format=serialization.PublicFormat.SubjectPublicKeyInfo, ) key_hash = hashlib.sha256(public_key_der).digest() - key_hash_b64 = base64.b64encode(key_hash).decode('utf-8') + key_hash_b64 = base64.b64encode(key_hash).decode("utf-8") logger.info(f"Public key DER length: {len(public_key_der)} bytes") logger.info(f"Public key SHA256 hash (hex): {key_hash.hex()}") @@ -93,28 +91,30 @@ def register_with_nitriding(self): url = f"{NITRIDING_BASE_URL}/enclave/hash" req = urllib.request.Request( - url, - data=key_hash_b64.encode('utf-8'), - method='POST' + url, data=key_hash_b64.encode("utf-8"), method="POST" ) response = urllib.request.urlopen(req, timeout=5) - response_body = response.read().decode('utf-8') + response_body = response.read().decode("utf-8") if response.getcode() == 200: logger.info("Successfully registered public key hash with nitriding") logger.info(f"Response: {response_body}") return True else: - logger.error(f"Failed to register public key hash: HTTP {response.getcode()}") + logger.error( + f"Failed to register public key hash: HTTP {response.getcode()}" + ) return False except urllib.error.HTTPError as e: - error_body = e.read().decode('utf-8') if e.fp else "No error body" + error_body = e.read().decode("utf-8") if e.fp else "No error body" logger.error(f"HTTP Error {e.code}: {e.reason} - {error_body}") return False except Exception as e: - logger.warning(f"Could not register with nitriding (may not be in TEE): {e}") + logger.warning( + f"Could not register with nitriding (may not be in TEE): {e}" + ) return False def sign_data(self, data: bytes) -> str: @@ -131,11 +131,11 @@ def sign_data(self, data: bytes) -> str: data, padding.PSS( mgf=padding.MGF1(hashes.SHA256()), - salt_length=hashes.SHA256().digest_size # 32 bytes + salt_length=hashes.SHA256().digest_size, # 32 bytes ), - hashes.SHA256() + hashes.SHA256(), ) - return base64.b64encode(signature).decode('utf-8') + return base64.b64encode(signature).decode("utf-8") def get_public_key(self) -> str: """Return public key in PEM format.""" @@ -159,9 +159,9 @@ def get_attestation_document(self) -> dict: "enclave_info": { "platform": "aws-nitro", "instance_type": "tee-enabled", - "version": "1.0.0" + "version": "1.0.0", }, - "measurements": None # Would contain PCR values in real deployment + "measurements": None, # Would contain PCR values in real deployment } @@ -189,10 +189,10 @@ def compute_tee_msg_hash( Returns (msg_hash_bytes, input_hash_hex, output_hash_hex). """ - input_hash = keccak(request_bytes) - output_hash = keccak(response_content.encode('utf-8')) - msg_hash = keccak(input_hash + output_hash + timestamp.to_bytes(32, 'big')) - logger.debug(f"Compute TEE Message Hash:\n\tInput hash: {input_hash.hex()} \n\tOutput hash: {output_hash.hex()}\n\tmsg_hash: {msg_hash}\n\ttimestamp: {timestamp} hashed timestamp:{timestamp.to_bytes(32, 'big')}") + input_hash = keccak(request_bytes) + output_hash = keccak(response_content.encode("utf-8")) + msg_hash = keccak(input_hash + output_hash + timestamp.to_bytes(32, "big")) + return msg_hash, input_hash.hex(), output_hash.hex() diff --git a/tee_gateway/test/__init__.py b/tee_gateway/test/__init__.py index 8979056..6caa498 100644 --- a/tee_gateway/test/__init__.py +++ b/tee_gateway/test/__init__.py @@ -7,10 +7,9 @@ class BaseTestCase(TestCase): - def create_app(self): - logging.getLogger('connexion.operation').setLevel('ERROR') - app = connexion.App(__name__, specification_dir='../openapi/') + logging.getLogger("connexion.operation").setLevel("ERROR") + app = connexion.App(__name__, specification_dir="../openapi/") app.app.json_encoder = JSONEncoder - app.add_api('openapi.yaml', pythonic_params=True) + app.add_api("openapi.yaml", pythonic_params=True) return app.app diff --git a/tee_gateway/test/test_chat_controller.py b/tee_gateway/test/test_chat_controller.py index c9f05a3..11c03b0 100644 --- a/tee_gateway/test/test_chat_controller.py +++ b/tee_gateway/test/test_chat_controller.py @@ -15,21 +15,108 @@ def test_create_chat_completion(self): Creates a model response for the given chat conversation via HTTP backend. Tests the HTTP-based chat completion endpoint that forwards requests to the TEE server. """ - body = {"reasoning_effort":"medium","top_logprobs":2,"metadata":{"key":"metadata"},"logit_bias":{"key":6},"seed":2147483647,"functions":[{"name":"name","description":"description","parameters":{"key":""}},{"name":"name","description":"description","parameters":{"key":""}},{"name":"name","description":"description","parameters":{"key":""}},{"name":"name","description":"description","parameters":{"key":""}},{"name":"name","description":"description","parameters":{"key":""}}],"function_call":"none","presence_penalty":-1.079145645226094,"tools":[{"function":{"name":"name","description":"description","strict":False,"parameters":{"key":""}},"type":"function"},{"function":{"name":"name","description":"description","strict":False,"parameters":{"key":""}},"type":"function"}],"logprobs":False,"top_p":1,"max_completion_tokens":5,"frequency_penalty":-1.6796687238155954,"modalities":["text","text"],"response_format":{"type":"text"},"stream":False,"temperature":1,"tool_choice":"none","model":"gpt-4o","service_tier":"auto","audio":{"voice":"alloy","format":"wav"},"max_tokens":5,"store":False,"n":1,"stop":"CreateChatCompletionRequest_stop","parallel_tool_calls":True,"prediction":{"type":"content","content":"PredictionContent_content"},"messages":[{"role":"developer","name":"name","content":"ChatCompletionRequestDeveloperMessage_content"},{"role":"developer","name":"name","content":"ChatCompletionRequestDeveloperMessage_content"}],"stream_options":{"include_usage":True},"user":"user-1234"} - headers = { - 'Accept': 'application/json', - 'Content-Type': 'application/json', - 'Authorization': 'Bearer special-key', + body = { + "reasoning_effort": "medium", + "top_logprobs": 2, + "metadata": {"key": "metadata"}, + "logit_bias": {"key": 6}, + "seed": 2147483647, + "functions": [ + { + "name": "name", + "description": "description", + "parameters": {"key": ""}, + }, + { + "name": "name", + "description": "description", + "parameters": {"key": ""}, + }, + { + "name": "name", + "description": "description", + "parameters": {"key": ""}, + }, + { + "name": "name", + "description": "description", + "parameters": {"key": ""}, + }, + { + "name": "name", + "description": "description", + "parameters": {"key": ""}, + }, + ], + "function_call": "none", + "presence_penalty": -1.079145645226094, + "tools": [ + { + "function": { + "name": "name", + "description": "description", + "strict": False, + "parameters": {"key": ""}, + }, + "type": "function", + }, + { + "function": { + "name": "name", + "description": "description", + "strict": False, + "parameters": {"key": ""}, + }, + "type": "function", + }, + ], + "logprobs": False, + "top_p": 1, + "max_completion_tokens": 5, + "frequency_penalty": -1.6796687238155954, + "modalities": ["text", "text"], + "response_format": {"type": "text"}, + "stream": False, + "temperature": 1, + "tool_choice": "none", + "model": "gpt-4o", + "service_tier": "auto", + "audio": {"voice": "alloy", "format": "wav"}, + "max_tokens": 5, + "store": False, + "n": 1, + "stop": "CreateChatCompletionRequest_stop", + "parallel_tool_calls": True, + "prediction": {"type": "content", "content": "PredictionContent_content"}, + "messages": [ + { + "role": "developer", + "name": "name", + "content": "ChatCompletionRequestDeveloperMessage_content", + }, + { + "role": "developer", + "name": "name", + "content": "ChatCompletionRequestDeveloperMessage_content", + }, + ], + "stream_options": {"include_usage": True}, + "user": "user-1234", + } + headers = { + "Accept": "application/json", + "Content-Type": "application/json", + "Authorization": "Bearer special-key", } response = self.client.open( - '/v1/chat/completions', - method='POST', + "/v1/chat/completions", + method="POST", headers=headers, data=json.dumps(body), - content_type='application/json') - self.assert200(response, - 'Response body is : ' + response.data.decode('utf-8')) + content_type="application/json", + ) + self.assert200(response, "Response body is : " + response.data.decode("utf-8")) -if __name__ == '__main__': - unittest.main() \ No newline at end of file +if __name__ == "__main__": + unittest.main() diff --git a/tee_gateway/test/test_completions_controller.py b/tee_gateway/test/test_completions_controller.py index b0c1b87..5591c06 100644 --- a/tee_gateway/test/test_completions_controller.py +++ b/tee_gateway/test/test_completions_controller.py @@ -15,21 +15,40 @@ def test_create_completion(self): Creates a completion for the provided prompt and parameters via HTTP backend. Tests the HTTP-based completion endpoint that forwards requests to the TEE server. """ - body = {"logit_bias":{"key":1},"seed":-2147483648,"max_tokens":16,"presence_penalty":0.25495066265333133,"echo":False,"suffix":"test.","n":1,"logprobs":2,"top_p":1,"frequency_frequency":0.4109824732281613,"best_of":1,"stop":"\n","stream":False,"temperature":1,"model":"CreateCompletionRequest_model","stream_options":{"include_usage":True},"prompt":"This is a test.","user":"user-1234"} - headers = { - 'Accept': 'application/json', - 'Content-Type': 'application/json', - 'Authorization': 'Bearer special-key', + body = { + "logit_bias": {"key": 1}, + "seed": -2147483648, + "max_tokens": 16, + "presence_penalty": 0.25495066265333133, + "echo": False, + "suffix": "test.", + "n": 1, + "logprobs": 2, + "top_p": 1, + "frequency_frequency": 0.4109824732281613, + "best_of": 1, + "stop": "\n", + "stream": False, + "temperature": 1, + "model": "CreateCompletionRequest_model", + "stream_options": {"include_usage": True}, + "prompt": "This is a test.", + "user": "user-1234", + } + headers = { + "Accept": "application/json", + "Content-Type": "application/json", + "Authorization": "Bearer special-key", } response = self.client.open( - '/v1/completions', - method='POST', + "/v1/completions", + method="POST", headers=headers, data=json.dumps(body), - content_type='application/json') - self.assert200(response, - 'Response body is : ' + response.data.decode('utf-8')) + content_type="application/json", + ) + self.assert200(response, "Response body is : " + response.data.decode("utf-8")) -if __name__ == '__main__': - unittest.main() \ No newline at end of file +if __name__ == "__main__": + unittest.main() diff --git a/tee_gateway/test/test_tool_forwarding.py b/tee_gateway/test/test_tool_forwarding.py index 71ab269..e5d4b3d 100644 --- a/tee_gateway/test/test_tool_forwarding.py +++ b/tee_gateway/test/test_tool_forwarding.py @@ -2,8 +2,8 @@ from unittest.mock import patch, Mock from tee_gateway.controllers.chat_controller import ( - parse_chat_request, - parse_message, + _parse_chat_request as parse_chat_request, + _parse_message as parse_message, create_chat_completion, ) from tee_gateway.models import ( @@ -20,7 +20,7 @@ def test_parse_tool_message(self): message_dict = { "role": "tool", "content": "The weather is sunny", - "tool_call_id": "call_123" + "tool_call_id": "call_123", } result = parse_message(message_dict) self.assertIsInstance(result, ChatCompletionRequestToolMessage) @@ -33,7 +33,7 @@ def test_parse_function_message(self): message_dict = { "role": "function", "content": "Result from function", - "name": "get_weather" + "name": "get_weather", } result = parse_message(message_dict) self.assertIsInstance(result, ChatCompletionRequestFunctionMessage) @@ -46,15 +46,17 @@ def test_parse_chat_request_with_tools(self): request_dict = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get the weather", - "parameters": {"type": "object", "properties": {}} + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get the weather", + "parameters": {"type": "object", "properties": {}}, + }, } - }], - "tool_choice": "auto" + ], + "tool_choice": "auto", } result = parse_chat_request(request_dict) self.assertEqual(result.model, "gpt-4") @@ -62,8 +64,8 @@ def test_parse_chat_request_with_tools(self): self.assertEqual(len(result.tools), 1) self.assertEqual(result.tool_choice, "auto") - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") def test_tools_forwarded_to_backend(self, mock_connexion, mock_http_session): """Test that tools are forwarded to HTTP backend""" # Setup mock request @@ -72,17 +74,22 @@ def test_tools_forwarded_to_backend(self, mock_connexion, mock_http_session): mock_connexion.request.get_json.return_value = { "model": "gpt-4", "messages": [{"role": "user", "content": "What is the weather?"}], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a location", - "parameters": {"type": "object", "properties": {"location": {"type": "string"}}}, - "strict": False + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": {"location": {"type": "string"}}, + }, + "strict": False, + }, } - }], + ], "tool_choice": "auto", - "stream": False + "stream": False, } # Setup mock HTTP response @@ -90,19 +97,12 @@ def test_tools_forwarded_to_backend(self, mock_connexion, mock_http_session): mock_response.status_code = 200 mock_response.json.return_value = { "finish_reason": "stop", - "message": { - "role": "assistant", - "content": "The weather is sunny." - }, + "message": {"role": "assistant", "content": "The weather is sunny."}, "model": "gpt-4", - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15 - }, + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, "timestamp": "2025-01-26T12:00:00Z", "signature": "test_signature", - "request_hash": "test_hash" + "request_hash": "test_hash", } mock_http_session.post.return_value = mock_response @@ -112,22 +112,24 @@ def test_tools_forwarded_to_backend(self, mock_connexion, mock_http_session): # Verify HTTP POST was called mock_http_session.post.assert_called_once() call_kwargs = mock_http_session.post.call_args[1] - + # Check that tools were included in the request - request_json = call_kwargs['json'] - self.assertIn('tools', request_json) - self.assertEqual(len(request_json['tools']), 1) - self.assertEqual(request_json['tools'][0]['type'], "function") - self.assertEqual(request_json['tools'][0]['function']['name'], "get_weather") - self.assertEqual(request_json['tool_choice'], "auto") - - # Verify response structure - self.assertIn('choices', result) - self.assertEqual(len(result['choices']), 1) + request_json = call_kwargs["json"] + self.assertIn("tools", request_json) + self.assertEqual(len(request_json["tools"]), 1) + self.assertEqual(request_json["tools"][0]["type"], "function") + self.assertEqual(request_json["tools"][0]["function"]["name"], "get_weather") + self.assertEqual(request_json["tool_choice"], "auto") - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') - def test_tool_calls_extracted_from_response(self, mock_connexion, mock_http_session): + # Verify response structure + self.assertIn("choices", result) + self.assertEqual(len(result["choices"]), 1) + + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") + def test_tool_calls_extracted_from_response( + self, mock_connexion, mock_http_session + ): """Test that tool_calls are extracted from HTTP response""" # Setup mock request mock_connexion.request.is_json = True @@ -135,7 +137,7 @@ def test_tool_calls_extracted_from_response(self, mock_connexion, mock_http_sess mock_connexion.request.get_json.return_value = { "model": "gpt-4", "messages": [{"role": "user", "content": "What is the weather?"}], - "stream": False + "stream": False, } # Setup mock HTTP response with tool_calls @@ -146,24 +148,22 @@ def test_tool_calls_extracted_from_response(self, mock_connexion, mock_http_sess "message": { "role": "assistant", "content": "", - "tool_calls": [{ - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "San Francisco"}' + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "San Francisco"}', + }, } - }] + ], }, "model": "gpt-4", - "usage": { - "prompt_tokens": 10, - "completion_tokens": 5, - "total_tokens": 15 - }, + "usage": {"prompt_tokens": 10, "completion_tokens": 5, "total_tokens": 15}, "timestamp": "2025-01-26T12:00:00Z", "signature": "test_signature", - "request_hash": "test_hash" + "request_hash": "test_hash", } mock_http_session.post.return_value = mock_response @@ -171,7 +171,7 @@ def test_tool_calls_extracted_from_response(self, mock_connexion, mock_http_sess response = create_chat_completion(None) # Verify tool_calls are in the response - self.assertIn('choices', response) + self.assertIn("choices", response) self.assertEqual(len(response["choices"]), 1) message = response["choices"][0]["message"] self.assertIn("tool_calls", message) @@ -179,10 +179,13 @@ def test_tool_calls_extracted_from_response(self, mock_connexion, mock_http_sess self.assertEqual(message["tool_calls"][0]["id"], "call_abc123") self.assertEqual(message["tool_calls"][0]["type"], "function") self.assertEqual(message["tool_calls"][0]["function"]["name"], "get_weather") - self.assertEqual(message["tool_calls"][0]["function"]["arguments"], '{"location": "San Francisco"}') + self.assertEqual( + message["tool_calls"][0]["function"]["arguments"], + '{"location": "San Francisco"}', + ) - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") def test_tool_message_forwarded_to_backend(self, mock_connexion, mock_http_session): """Test that tool messages in the conversation are forwarded to HTTP backend""" # Setup mock request with tool message @@ -195,22 +198,24 @@ def test_tool_message_forwarded_to_backend(self, mock_connexion, mock_http_sessi { "role": "assistant", "content": None, - "tool_calls": [{ - "id": "call_abc123", - "type": "function", - "function": { - "name": "get_weather", - "arguments": '{"location": "San Francisco"}' + "tool_calls": [ + { + "id": "call_abc123", + "type": "function", + "function": { + "name": "get_weather", + "arguments": '{"location": "San Francisco"}', + }, } - }] + ], }, { "role": "tool", "content": '{"temperature": 72, "condition": "sunny"}', - "tool_call_id": "call_abc123" - } + "tool_call_id": "call_abc123", + }, ], - "stream": False + "stream": False, } # Setup mock HTTP response @@ -220,17 +225,13 @@ def test_tool_message_forwarded_to_backend(self, mock_connexion, mock_http_sessi "finish_reason": "stop", "message": { "role": "assistant", - "content": "The weather in San Francisco is 72°F and sunny." + "content": "The weather in San Francisco is 72°F and sunny.", }, "model": "gpt-4", - "usage": { - "prompt_tokens": 20, - "completion_tokens": 12, - "total_tokens": 32 - }, + "usage": {"prompt_tokens": 20, "completion_tokens": 12, "total_tokens": 32}, "timestamp": "2025-01-26T12:00:00Z", "signature": "test_signature", - "request_hash": "test_hash" + "request_hash": "test_hash", } mock_http_session.post.return_value = mock_response @@ -240,33 +241,37 @@ def test_tool_message_forwarded_to_backend(self, mock_connexion, mock_http_sessi # Verify HTTP POST was called with tool message mock_http_session.post.assert_called_once() call_kwargs = mock_http_session.post.call_args[1] - - request_json = call_kwargs['json'] - messages = request_json['messages'] + + request_json = call_kwargs["json"] + messages = request_json["messages"] # Should have 3 messages: user, assistant with tool_calls, tool self.assertEqual(len(messages), 3) # Check tool message was converted tool_msg = messages[2] - self.assertEqual(tool_msg['role'], "tool") - self.assertEqual(tool_msg['content'], '{"temperature": 72, "condition": "sunny"}') - self.assertEqual(tool_msg['tool_call_id'], "call_abc123") + self.assertEqual(tool_msg["role"], "tool") + self.assertEqual( + tool_msg["content"], '{"temperature": 72, "condition": "sunny"}' + ) + self.assertEqual(tool_msg["tool_call_id"], "call_abc123") # Check assistant message has tool_calls assistant_msg = messages[1] - self.assertEqual(assistant_msg['role'], "assistant") - self.assertIn('tool_calls', assistant_msg) - self.assertEqual(len(assistant_msg['tool_calls']), 1) - self.assertEqual(assistant_msg['tool_calls'][0]['id'], "call_abc123") - self.assertEqual(assistant_msg['tool_calls'][0]['function']['name'], "get_weather") + self.assertEqual(assistant_msg["role"], "assistant") + self.assertIn("tool_calls", assistant_msg) + self.assertEqual(len(assistant_msg["tool_calls"]), 1) + self.assertEqual(assistant_msg["tool_calls"][0]["id"], "call_abc123") + self.assertEqual( + assistant_msg["tool_calls"][0]["function"]["name"], "get_weather" + ) # Verify response - self.assertIn('choices', result) - self.assertEqual(result['choices'][0]['finish_reason'], "stop") + self.assertIn("choices", result) + self.assertEqual(result["choices"][0]["finish_reason"], "stop") - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") def test_payment_header_forwarded(self, mock_connexion, mock_http_session): """Test that X-PAYMENT header is forwarded to backend""" # Setup mock request with payment header @@ -275,7 +280,7 @@ def test_payment_header_forwarded(self, mock_connexion, mock_http_session): mock_connexion.request.get_json.return_value = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "stream": False + "stream": False, } # Setup mock HTTP response @@ -283,16 +288,9 @@ def test_payment_header_forwarded(self, mock_connexion, mock_http_session): mock_response.status_code = 200 mock_response.json.return_value = { "finish_reason": "stop", - "message": { - "role": "assistant", - "content": "Hello!" - }, + "message": {"role": "assistant", "content": "Hello!"}, "model": "gpt-4", - "usage": { - "prompt_tokens": 5, - "completion_tokens": 2, - "total_tokens": 7 - } + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, } mock_http_session.post.return_value = mock_response @@ -302,13 +300,13 @@ def test_payment_header_forwarded(self, mock_connexion, mock_http_session): # Verify X-PAYMENT header was included mock_http_session.post.assert_called_once() call_kwargs = mock_http_session.post.call_args[1] - - self.assertIn('headers', call_kwargs) - self.assertIn('X-PAYMENT', call_kwargs['headers']) - self.assertEqual(call_kwargs['headers']['X-PAYMENT'], "payment_token_123") - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') + self.assertIn("headers", call_kwargs) + self.assertIn("X-PAYMENT", call_kwargs["headers"]) + self.assertEqual(call_kwargs["headers"]["X-PAYMENT"], "payment_token_123") + + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") def test_tee_metadata_preserved(self, mock_connexion, mock_http_session): """Test that TEE metadata is preserved in response""" # Setup mock request @@ -317,7 +315,7 @@ def test_tee_metadata_preserved(self, mock_connexion, mock_http_session): mock_connexion.request.get_json.return_value = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "stream": False + "stream": False, } # Setup mock HTTP response with TEE metadata @@ -325,19 +323,12 @@ def test_tee_metadata_preserved(self, mock_connexion, mock_http_session): mock_response.status_code = 200 mock_response.json.return_value = { "finish_reason": "stop", - "message": { - "role": "assistant", - "content": "Hello!" - }, + "message": {"role": "assistant", "content": "Hello!"}, "model": "gpt-4", - "usage": { - "prompt_tokens": 5, - "completion_tokens": 2, - "total_tokens": 7 - }, + "usage": {"prompt_tokens": 5, "completion_tokens": 2, "total_tokens": 7}, "timestamp": "2025-01-26T12:00:00Z", "signature": "tee_signature_abc123", - "request_hash": "hash_def456" + "request_hash": "hash_def456", } mock_http_session.post.return_value = mock_response @@ -345,20 +336,20 @@ def test_tee_metadata_preserved(self, mock_connexion, mock_http_session): result = create_chat_completion(None) # Verify basic response structure - self.assertIn('choices', result) - self.assertIn('model', result) - self.assertEqual(result['model'], "gpt-4") - + self.assertIn("choices", result) + self.assertIn("model", result) + self.assertEqual(result["model"], "gpt-4") + # Verify TEE metadata is preserved in response - self.assertIn('tee_signature', result) - self.assertEqual(result['tee_signature'], "tee_signature_abc123") - self.assertIn('tee_request_hash', result) - self.assertEqual(result['tee_request_hash'], "hash_def456") - self.assertIn('tee_timestamp', result) - self.assertEqual(result['tee_timestamp'], "2025-01-26T12:00:00Z") - - @patch('tee_gateway.controllers.chat_controller.http_session') - @patch('tee_gateway.controllers.chat_controller.connexion') + self.assertIn("tee_signature", result) + self.assertEqual(result["tee_signature"], "tee_signature_abc123") + self.assertIn("tee_request_hash", result) + self.assertEqual(result["tee_request_hash"], "hash_def456") + self.assertIn("tee_timestamp", result) + self.assertEqual(result["tee_timestamp"], "2025-01-26T12:00:00Z") + + @patch("tee_gateway.controllers.chat_controller.http_session") + @patch("tee_gateway.controllers.chat_controller.connexion") def test_http_error_handling(self, mock_connexion, mock_http_session): """Test that HTTP errors are handled properly""" # Setup mock request @@ -367,22 +358,25 @@ def test_http_error_handling(self, mock_connexion, mock_http_session): mock_connexion.request.get_json.return_value = { "model": "gpt-4", "messages": [{"role": "user", "content": "Hello"}], - "stream": False + "stream": False, } # Setup mock to raise HTTP error import requests - mock_http_session.post.side_effect = requests.exceptions.RequestException("Connection failed") + + mock_http_session.post.side_effect = requests.exceptions.RequestException( + "Connection failed" + ) # Call the function result, status_code = create_chat_completion(None) # Verify error response self.assertEqual(status_code, 500) - self.assertIn('error', result) - self.assertEqual(result['error'], "Backend request failed") - self.assertIn('details', result) + self.assertIn("error", result) + self.assertEqual(result["error"], "Backend request failed") + self.assertIn("details", result) -if __name__ == '__main__': - unittest.main() \ No newline at end of file +if __name__ == "__main__": + unittest.main() diff --git a/tee_gateway/typing_utils.py b/tee_gateway/typing_utils.py index 74e3c91..c905b1d 100644 --- a/tee_gateway/typing_utils.py +++ b/tee_gateway/typing_utils.py @@ -4,27 +4,27 @@ import typing def is_generic(klass): - """ Determine whether klass is a generic class """ + """Determine whether klass is a generic class""" return type(klass) == typing.GenericMeta def is_dict(klass): - """ Determine whether klass is a Dict """ + """Determine whether klass is a Dict""" return klass.__extra__ == dict def is_list(klass): - """ Determine whether klass is a List """ + """Determine whether klass is a List""" return klass.__extra__ == list else: def is_generic(klass): - """ Determine whether klass is a generic class """ - return hasattr(klass, '__origin__') + """Determine whether klass is a generic class""" + return hasattr(klass, "__origin__") def is_dict(klass): - """ Determine whether klass is a Dict """ + """Determine whether klass is a Dict""" return klass.__origin__ == dict def is_list(klass): - """ Determine whether klass is a List """ + """Determine whether klass is a List""" return klass.__origin__ == list diff --git a/tee_gateway/util.py b/tee_gateway/util.py index 63801cd..d6336ce 100644 --- a/tee_gateway/util.py +++ b/tee_gateway/util.py @@ -73,10 +73,11 @@ def deserialize_date(string): :rtype: date """ if string is None: - return None - + return None + try: from dateutil.parser import parse + return parse(string).date() except ImportError: return string @@ -93,10 +94,11 @@ def deserialize_datetime(string): :rtype: datetime """ if string is None: - return None - + return None + try: from dateutil.parser import parse + return parse(string) except ImportError: return string @@ -116,9 +118,11 @@ def deserialize_model(data, klass): return data for attr, attr_type in instance.openapi_types.items(): - if data is not None \ - and instance.attribute_map[attr] in data \ - and isinstance(data, (list, dict)): + if ( + data is not None + and instance.attribute_map[attr] in data + and isinstance(data, (list, dict)) + ): value = data[instance.attribute_map[attr]] setattr(instance, attr, _deserialize(value, attr_type)) @@ -135,8 +139,7 @@ def _deserialize_list(data, boxed_type): :return: deserialized list. :rtype: list """ - return [_deserialize(sub_data, boxed_type) - for sub_data in data] + return [_deserialize(sub_data, boxed_type) for sub_data in data] def _deserialize_dict(data, boxed_type): @@ -149,14 +152,15 @@ def _deserialize_dict(data, boxed_type): :return: deserialized dict. :rtype: dict """ - return {k: _deserialize(v, boxed_type) - for k, v in data.items() } + return {k: _deserialize(v, boxed_type) for k, v in data.items()} + from tee_gateway.definitions import ( # noqa: E402 ASSET_DECIMALS_BY_ADDRESS, DEFAULT_ASSET_DECIMALS, ) from tee_gateway.model_registry import get_model_config # noqa: E402 + TOKEN_A_PRICE_CACHE_TTL_SECONDS = 60 _token_price_cache: dict[str, Any] = { @@ -183,7 +187,10 @@ def get_token_a_price_usd() -> Decimal: with _token_price_lock: cached_value = _token_price_cache.get("value") cached_at = float(_token_price_cache.get("updated_at") or 0.0) - if isinstance(cached_value, Decimal) and (now - cached_at) < TOKEN_A_PRICE_CACHE_TTL_SECONDS: + if ( + isinstance(cached_value, Decimal) + and (now - cached_at) < TOKEN_A_PRICE_CACHE_TTL_SECONDS + ): return cached_value value = _fetch_token_a_price_usd_mock() @@ -229,7 +236,9 @@ def _normalize_model_name(model: str | None) -> str | None: return str(model).strip().lower() -def _extract_usage_tokens(response_json: dict[str, Any] | None) -> tuple[int, int] | None: +def _extract_usage_tokens( + response_json: dict[str, Any] | None, +) -> tuple[int, int] | None: if not isinstance(response_json, dict): return None usage = response_json.get("usage") @@ -277,7 +286,9 @@ def dynamic_session_cost_calculator(context: dict[str, Any]) -> int: response_json = context.get("response_json") if not isinstance(request_json, dict) or not isinstance(response_json, dict): - raise ValueError("dynamic_session_cost_calculator requires both request_json and response_json") + raise ValueError( + "dynamic_session_cost_calculator requires both request_json and response_json" + ) model = _extract_model_from_context(request_json, response_json) if not model: @@ -296,15 +307,21 @@ def dynamic_session_cost_calculator(context: dict[str, Any]) -> int: input_rate = cfg.input_price_usd output_rate = cfg.output_price_usd - total_usd = (Decimal(input_tokens) * input_rate) + (Decimal(output_tokens) * output_rate) + total_usd = (Decimal(input_tokens) * input_rate) + ( + Decimal(output_tokens) * output_rate + ) token_price_usd = get_token_a_price_usd() if token_price_usd <= 0: raise ValueError(f"Token A price is non-positive: {token_price_usd}") token_amount = total_usd / token_price_usd - decimals = _extract_asset_decimals_from_requirements(context.get("payment_requirements")) + decimals = _extract_asset_decimals_from_requirements( + context.get("payment_requirements") + ) scale = Decimal(10) ** decimals - cost_smallest_units = int((token_amount * scale).to_integral_value(rounding=ROUND_CEILING)) + cost_smallest_units = int( + (token_amount * scale).to_integral_value(rounding=ROUND_CEILING) + ) logger.info( "DYNAMIC_SESSION_COST model=%s input_tokens=%d output_tokens=%d total_usd=%s token_price_usd=%s decimals=%d cost=%d", diff --git a/tests/test_server.py b/tests/test_server.py index 06036c6..faa77d4 100644 --- a/tests/test_server.py +++ b/tests/test_server.py @@ -28,29 +28,36 @@ from fastapi.testclient import TestClient # Mock the TEE key manager registration before importing the app -with patch('urllib.request.urlopen'): - from server import ( - app - ) +with patch("urllib.request.urlopen"): + from server import app client = TestClient(app) class MockAIMessage: """Mock LangChain AIMessage for testing""" - def __init__(self, content: str, tool_calls: List[Dict] = None, usage_metadata: Dict = None): + + def __init__( + self, content: str, tool_calls: List[Dict] = None, usage_metadata: Dict = None + ): self.content = content self.tool_calls = tool_calls or [] self.usage_metadata = usage_metadata or { "input_tokens": 10, "output_tokens": 20, - "total_tokens": 30 + "total_tokens": 30, } class MockStreamChunk: """Mock streaming chunk for testing""" - def __init__(self, content: str = "", tool_call_chunks: List[Dict] = None, usage_metadata: Dict = None): + + def __init__( + self, + content: str = "", + tool_call_chunks: List[Dict] = None, + usage_metadata: Dict = None, + ): self.content = content self.tool_call_chunks = tool_call_chunks or [] self.usage_metadata = usage_metadata @@ -61,11 +68,13 @@ def __init__(self, content: str = "", tool_call_chunks: List[Dict] = None, usage def mock_openai_model(): """Mock OpenAI chat model""" mock_model = MagicMock() - mock_model.invoke = MagicMock(return_value=MockAIMessage( - content="Hello from OpenAI!", - usage_metadata={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15} - )) - + mock_model.invoke = MagicMock( + return_value=MockAIMessage( + content="Hello from OpenAI!", + usage_metadata={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15}, + ) + ) + async def mock_astream(messages): """Mock async streaming""" # Yield content chunks @@ -73,8 +82,10 @@ async def mock_astream(messages): yield MockStreamChunk(content="from ") yield MockStreamChunk(content="OpenAI!") # Yield final usage - yield MockStreamChunk(usage_metadata={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15}) - + yield MockStreamChunk( + usage_metadata={"input_tokens": 5, "output_tokens": 10, "total_tokens": 15} + ) + mock_model.astream = mock_astream mock_model.bind_tools = MagicMock(return_value=mock_model) return mock_model @@ -84,17 +95,21 @@ async def mock_astream(messages): def mock_anthropic_model(): """Mock Anthropic chat model""" mock_model = MagicMock() - mock_model.invoke = MagicMock(return_value=MockAIMessage( - content="Hello from Anthropic!", - usage_metadata={"input_tokens": 8, "output_tokens": 12, "total_tokens": 20} - )) - + mock_model.invoke = MagicMock( + return_value=MockAIMessage( + content="Hello from Anthropic!", + usage_metadata={"input_tokens": 8, "output_tokens": 12, "total_tokens": 20}, + ) + ) + async def mock_astream(messages): yield MockStreamChunk(content="Hello ") yield MockStreamChunk(content="from ") yield MockStreamChunk(content="Anthropic!") - yield MockStreamChunk(usage_metadata={"input_tokens": 8, "output_tokens": 12, "total_tokens": 20}) - + yield MockStreamChunk( + usage_metadata={"input_tokens": 8, "output_tokens": 12, "total_tokens": 20} + ) + mock_model.astream = mock_astream mock_model.bind_tools = MagicMock(return_value=mock_model) return mock_model @@ -104,17 +119,21 @@ async def mock_astream(messages): def mock_google_model(): """Mock Google chat model""" mock_model = MagicMock() - mock_model.invoke = MagicMock(return_value=MockAIMessage( - content="Hello from Google!", - usage_metadata={"input_tokens": 6, "output_tokens": 11, "total_tokens": 17} - )) - + mock_model.invoke = MagicMock( + return_value=MockAIMessage( + content="Hello from Google!", + usage_metadata={"input_tokens": 6, "output_tokens": 11, "total_tokens": 17}, + ) + ) + async def mock_astream(messages): yield MockStreamChunk(content="Hello ") yield MockStreamChunk(content="from ") yield MockStreamChunk(content="Google!") - yield MockStreamChunk(usage_metadata={"input_tokens": 6, "output_tokens": 11, "total_tokens": 17}) - + yield MockStreamChunk( + usage_metadata={"input_tokens": 6, "output_tokens": 11, "total_tokens": 17} + ) + mock_model.astream = mock_astream mock_model.bind_tools = MagicMock(return_value=mock_model) return mock_model @@ -124,17 +143,21 @@ async def mock_astream(messages): def mock_xai_model(): """Mock xAI chat model""" mock_model = MagicMock() - mock_model.invoke = MagicMock(return_value=MockAIMessage( - content="Hello from xAI!", - usage_metadata={"input_tokens": 7, "output_tokens": 13, "total_tokens": 20} - )) - + mock_model.invoke = MagicMock( + return_value=MockAIMessage( + content="Hello from xAI!", + usage_metadata={"input_tokens": 7, "output_tokens": 13, "total_tokens": 20}, + ) + ) + async def mock_astream(messages): yield MockStreamChunk(content="Hello ") yield MockStreamChunk(content="from ") yield MockStreamChunk(content="xAI!") - yield MockStreamChunk(usage_metadata={"input_tokens": 7, "output_tokens": 13, "total_tokens": 20}) - + yield MockStreamChunk( + usage_metadata={"input_tokens": 7, "output_tokens": 13, "total_tokens": 20} + ) + mock_model.astream = mock_astream mock_model.bind_tools = MagicMock(return_value=mock_model) return mock_model @@ -144,29 +167,43 @@ async def mock_astream(messages): def mock_tool_call_model(): """Mock model that returns tool calls""" mock_model = MagicMock() - mock_model.invoke = MagicMock(return_value=MockAIMessage( - content="", - tool_calls=[{ - "id": "call_123", - "name": "get_weather", - "args": {"location": "San Francisco", "unit": "celsius"}, - "type": "function" - }], - usage_metadata={"input_tokens": 15, "output_tokens": 25, "total_tokens": 40} - )) - + mock_model.invoke = MagicMock( + return_value=MockAIMessage( + content="", + tool_calls=[ + { + "id": "call_123", + "name": "get_weather", + "args": {"location": "San Francisco", "unit": "celsius"}, + "type": "function", + } + ], + usage_metadata={ + "input_tokens": 15, + "output_tokens": 25, + "total_tokens": 40, + }, + ) + ) + async def mock_astream_tools(messages): # Stream tool call chunks - yield MockStreamChunk(tool_call_chunks=[{ - "index": 0, - "id": "call_123", - "name": "get_weather", - "args": '{"location": "San Francisco", "unit": "celsius"}', - "type": "function" - }]) + yield MockStreamChunk( + tool_call_chunks=[ + { + "index": 0, + "id": "call_123", + "name": "get_weather", + "args": '{"location": "San Francisco", "unit": "celsius"}', + "type": "function", + } + ] + ) # Final usage - yield MockStreamChunk(usage_metadata={"input_tokens": 15, "output_tokens": 25, "total_tokens": 40}) - + yield MockStreamChunk( + usage_metadata={"input_tokens": 15, "output_tokens": 25, "total_tokens": 40} + ) + mock_model.astream = mock_astream_tools mock_model.bind_tools = MagicMock(return_value=mock_model) return mock_model @@ -204,7 +241,7 @@ def test_list_models(): data = response.json() assert "data" in data assert len(data["data"]) > 0 - + # Check for models from each provider model_ids = [m["id"] for m in data["data"]] assert any("gpt" in m for m in model_ids) # OpenAI @@ -214,18 +251,21 @@ def test_list_models(): # Completion Tests -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_completion_openai(mock_get_model, mock_openai_model): """Test OpenAI completion endpoint""" mock_get_model.return_value = mock_openai_model - - response = client.post("/v1/completions", json={ - "model": "gpt-4o", - "prompt": "Say hello", - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/completions", + json={ + "model": "gpt-4o", + "prompt": "Say hello", + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "completion" in data @@ -237,18 +277,21 @@ def test_completion_openai(mock_get_model, mock_openai_model): assert "request_hash" in data -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_completion_anthropic(mock_get_model, mock_anthropic_model): """Test Anthropic completion endpoint""" mock_get_model.return_value = mock_anthropic_model - - response = client.post("/v1/completions", json={ - "model": "claude-3.7-sonnet", - "prompt": "Say hello", - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/completions", + json={ + "model": "claude-3.7-sonnet", + "prompt": "Say hello", + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "completion" in data @@ -259,20 +302,21 @@ def test_completion_anthropic(mock_get_model, mock_anthropic_model): # Chat Completion Tests -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_completion_openai(mock_get_model, mock_openai_model): """Test OpenAI chat completion endpoint""" mock_get_model.return_value = mock_openai_model - - response = client.post("/v1/chat/completions", json={ - "model": "gpt-4o", - "messages": [ - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Say hello"}], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "message" in data @@ -285,21 +329,24 @@ def test_chat_completion_openai(mock_get_model, mock_openai_model): assert "signature" in data -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_completion_anthropic(mock_get_model, mock_anthropic_model): """Test Anthropic chat completion endpoint""" mock_get_model.return_value = mock_anthropic_model - - response = client.post("/v1/chat/completions", json={ - "model": "claude-4.0-sonnet", - "messages": [ - {"role": "system", "content": "You are helpful"}, - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "claude-4.0-sonnet", + "messages": [ + {"role": "system", "content": "You are helpful"}, + {"role": "user", "content": "Say hello"}, + ], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "message" in data @@ -309,20 +356,21 @@ def test_chat_completion_anthropic(mock_get_model, mock_anthropic_model): assert data["usage"]["total_tokens"] == 20 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_completion_google(mock_get_model, mock_google_model): """Test Google chat completion endpoint""" mock_get_model.return_value = mock_google_model - - response = client.post("/v1/chat/completions", json={ - "model": "gemini-2.5-flash-preview", - "messages": [ - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "gemini-2.5-flash-preview", + "messages": [{"role": "user", "content": "Say hello"}], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "message" in data @@ -332,20 +380,21 @@ def test_chat_completion_google(mock_get_model, mock_google_model): assert data["usage"]["total_tokens"] == 17 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_completion_xai(mock_get_model, mock_xai_model): """Test xAI chat completion endpoint""" mock_get_model.return_value = mock_xai_model - - response = client.post("/v1/chat/completions", json={ - "model": "grok-3-beta", - "messages": [ - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "grok-3-beta", + "messages": [{"role": "user", "content": "Say hello"}], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert "message" in data @@ -356,45 +405,53 @@ def test_chat_completion_xai(mock_get_model, mock_xai_model): # Tool Call Tests -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_tool_calls_openai(mock_get_model, mock_tool_call_model): """Test OpenAI chat with tool calls""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions", json={ - "model": "gpt-4o", - "messages": [ - {"role": "user", "content": "What's the weather in San Francisco?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a location", - "parameters": { - "type": "object", - "properties": { - "location": {"type": "string"}, - "unit": {"type": "string", "enum": ["celsius", "fahrenheit"]} + + response = client.post( + "/v1/chat/completions", + json={ + "model": "gpt-4o", + "messages": [ + {"role": "user", "content": "What's the weather in San Francisco?"} + ], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"}, + "unit": { + "type": "string", + "enum": ["celsius", "fahrenheit"], + }, + }, + "required": ["location"], + }, }, - "required": ["location"] } - } - }], - "temperature": 0.7, - "max_tokens": 100 - }) - + ], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 data = response.json() assert data["finish_reason"] == "tool_calls" assert "tool_calls" in data["message"] assert len(data["message"]["tool_calls"]) == 1 - + tool_call = data["message"]["tool_calls"][0] assert tool_call["type"] == "function" assert tool_call["function"]["name"] == "get_weather" - + # Parse arguments args = json.loads(tool_call["function"]["arguments"]) assert args["location"] == "San Francisco" @@ -402,33 +459,38 @@ def test_chat_tool_calls_openai(mock_get_model, mock_tool_call_model): assert data["usage"]["total_tokens"] == 40 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_tool_calls_anthropic(mock_get_model, mock_tool_call_model): """Test Anthropic chat with tool calls""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions", json={ - "model": "claude-4.0-sonnet", - "messages": [ - {"role": "user", "content": "What's the weather in San Francisco?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather for a location", - "parameters": { - "type": "object", - "properties": { - "location": {"type": "string"}, - "unit": {"type": "string"} - } + + response = client.post( + "/v1/chat/completions", + json={ + "model": "claude-4.0-sonnet", + "messages": [ + {"role": "user", "content": "What's the weather in San Francisco?"} + ], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather for a location", + "parameters": { + "type": "object", + "properties": { + "location": {"type": "string"}, + "unit": {"type": "string"}, + }, + }, + }, } - } - }], - "temperature": 0.7 - }) - + ], + "temperature": 0.7, + }, + ) + assert response.status_code == 200 data = response.json() assert data["finish_reason"] == "tool_calls" @@ -436,249 +498,262 @@ def test_chat_tool_calls_anthropic(mock_get_model, mock_tool_call_model): assert len(data["message"]["tool_calls"]) > 0 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_tool_calls_google(mock_get_model, mock_tool_call_model): """Test Google chat with tool calls""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions", json={ - "model": "gemini-2.5-flash-preview", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "gemini-2.5-flash-preview", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 data = response.json() assert data["finish_reason"] == "tool_calls" -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_tool_calls_xai(mock_get_model, mock_tool_call_model): """Test xAI chat with tool calls""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions", json={ - "model": "grok-3-beta", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions", + json={ + "model": "grok-3-beta", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 data = response.json() assert data["finish_reason"] == "tool_calls" # Streaming Tests -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_openai(mock_get_model, mock_openai_model): """Test OpenAI streaming chat completion""" mock_get_model.return_value = mock_openai_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "gpt-4o", - "messages": [ - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7, - "max_tokens": 100 - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "Say hello"}], + "temperature": 0.7, + "max_tokens": 100, + }, + ) + assert response.status_code == 200 assert response.headers["content-type"] == "text/event-stream; charset=utf-8" - + # Parse SSE stream content_chunks = [] usage_data = None finish_reason = None - - for line in response.text.split('\n'): - if line.startswith('data: '): + + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'content' in choice['delta']: - content_chunks.append(choice['delta']['content']) - if choice.get('finish_reason'): - finish_reason = choice['finish_reason'] - if 'usage' in chunk: - usage_data = chunk['usage'] + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "content" in choice["delta"]: + content_chunks.append(choice["delta"]["content"]) + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] + if "usage" in chunk: + usage_data = chunk["usage"] except json.JSONDecodeError: pass - - full_content = ''.join(content_chunks) + + full_content = "".join(content_chunks) assert "Hello from OpenAI!" in full_content assert finish_reason == "stop" assert usage_data is not None assert usage_data["total_tokens"] == 15 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_anthropic(mock_get_model, mock_anthropic_model): """Test Anthropic streaming chat completion""" mock_get_model.return_value = mock_anthropic_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "claude-4.0-sonnet", - "messages": [ - {"role": "user", "content": "Say hello"} - ], - "temperature": 0.7 - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "claude-4.0-sonnet", + "messages": [{"role": "user", "content": "Say hello"}], + "temperature": 0.7, + }, + ) + assert response.status_code == 200 - + content_chunks = [] - for line in response.text.split('\n'): - if line.startswith('data: '): + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'content' in choice['delta']: - content_chunks.append(choice['delta']['content']) + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "content" in choice["delta"]: + content_chunks.append(choice["delta"]["content"]) except json.JSONDecodeError: pass - - full_content = ''.join(content_chunks) + + full_content = "".join(content_chunks) assert "Hello from Anthropic!" in full_content -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_google(mock_get_model, mock_google_model): """Test Google streaming chat completion""" mock_get_model.return_value = mock_google_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "gemini-2.5-flash-preview", - "messages": [ - {"role": "user", "content": "Say hello"} - ] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "gemini-2.5-flash-preview", + "messages": [{"role": "user", "content": "Say hello"}], + }, + ) + assert response.status_code == 200 - + content_chunks = [] - for line in response.text.split('\n'): - if line.startswith('data: '): + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'content' in choice['delta']: - content_chunks.append(choice['delta']['content']) + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "content" in choice["delta"]: + content_chunks.append(choice["delta"]["content"]) except json.JSONDecodeError: pass - - full_content = ''.join(content_chunks) + + full_content = "".join(content_chunks) assert "Hello from Google!" in full_content -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_xai(mock_get_model, mock_xai_model): """Test xAI streaming chat completion""" mock_get_model.return_value = mock_xai_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "grok-3-beta", - "messages": [ - {"role": "user", "content": "Say hello"} - ] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "grok-3-beta", + "messages": [{"role": "user", "content": "Say hello"}], + }, + ) + assert response.status_code == 200 - + content_chunks = [] - for line in response.text.split('\n'): - if line.startswith('data: '): + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'content' in choice['delta']: - content_chunks.append(choice['delta']['content']) + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "content" in choice["delta"]: + content_chunks.append(choice["delta"]["content"]) except json.JSONDecodeError: pass - - full_content = ''.join(content_chunks) + + full_content = "".join(content_chunks) assert "Hello from xAI!" in full_content # Streaming with Tool Calls Tests -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_tool_calls_openai(mock_get_model, mock_tool_call_model): """Test OpenAI streaming with tool calls (buffered)""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "gpt-4o", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "gpt-4o", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 - + tool_calls = [] finish_reason = None usage_data = None - - for line in response.text.split('\n'): - if line.startswith('data: '): + + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'tool_calls' in choice['delta']: - tool_calls.extend(choice['delta']['tool_calls']) - if choice.get('finish_reason'): - finish_reason = choice['finish_reason'] - if 'usage' in chunk: - usage_data = chunk['usage'] + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "tool_calls" in choice["delta"]: + tool_calls.extend(choice["delta"]["tool_calls"]) + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] + if "usage" in chunk: + usage_data = chunk["usage"] except json.JSONDecodeError: pass - + # OpenAI buffers tool calls, so we should get complete tool calls assert len(tool_calls) > 0 assert finish_reason == "tool_calls" @@ -686,138 +761,147 @@ def test_chat_streaming_tool_calls_openai(mock_get_model, mock_tool_call_model): assert usage_data["total_tokens"] == 40 -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_tool_calls_anthropic(mock_get_model, mock_tool_call_model): """Test Anthropic streaming with tool calls (buffered)""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "claude-4.0-sonnet", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "claude-4.0-sonnet", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 - + finish_reason = None has_tool_calls = False - - for line in response.text.split('\n'): - if line.startswith('data: '): + + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'tool_calls' in choice['delta']: + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "tool_calls" in choice["delta"]: has_tool_calls = True - if choice.get('finish_reason'): - finish_reason = choice['finish_reason'] + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] except json.JSONDecodeError: pass - + assert has_tool_calls assert finish_reason == "tool_calls" -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_tool_calls_google(mock_get_model, mock_tool_call_model): """Test Google streaming with tool calls (not buffered)""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "gemini-2.5-flash-preview", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "gemini-2.5-flash-preview", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 - + finish_reason = None has_tool_calls = False - - for line in response.text.split('\n'): - if line.startswith('data: '): + + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'tool_calls' in choice['delta']: + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "tool_calls" in choice["delta"]: has_tool_calls = True - if choice.get('finish_reason'): - finish_reason = choice['finish_reason'] + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] except json.JSONDecodeError: pass - + # Google streams tool calls immediately assert has_tool_calls assert finish_reason == "tool_calls" -@patch('server.get_chat_model_cached') +@patch("server.get_chat_model_cached") def test_chat_streaming_tool_calls_xai(mock_get_model, mock_tool_call_model): """Test xAI streaming with tool calls (not buffered)""" mock_get_model.return_value = mock_tool_call_model - - response = client.post("/v1/chat/completions/stream", json={ - "model": "grok-3-beta", - "messages": [ - {"role": "user", "content": "What's the weather?"} - ], - "tools": [{ - "type": "function", - "function": { - "name": "get_weather", - "description": "Get weather", - "parameters": {"type": "object", "properties": {}} - } - }] - }) - + + response = client.post( + "/v1/chat/completions/stream", + json={ + "model": "grok-3-beta", + "messages": [{"role": "user", "content": "What's the weather?"}], + "tools": [ + { + "type": "function", + "function": { + "name": "get_weather", + "description": "Get weather", + "parameters": {"type": "object", "properties": {}}, + }, + } + ], + }, + ) + assert response.status_code == 200 - + finish_reason = None has_tool_calls = False - - for line in response.text.split('\n'): - if line.startswith('data: '): + + for line in response.text.split("\n"): + if line.startswith("data: "): data_str = line[6:] - if data_str == '[DONE]': + if data_str == "[DONE]": break try: chunk = json.loads(data_str) - if 'choices' in chunk and len(chunk['choices']) > 0: - choice = chunk['choices'][0] - if 'delta' in choice and 'tool_calls' in choice['delta']: + if "choices" in chunk and len(chunk["choices"]) > 0: + choice = chunk["choices"][0] + if "delta" in choice and "tool_calls" in choice["delta"]: has_tool_calls = True - if choice.get('finish_reason'): - finish_reason = choice['finish_reason'] + if choice.get("finish_reason"): + finish_reason = choice["finish_reason"] except json.JSONDecodeError: pass - + assert has_tool_calls assert finish_reason == "tool_calls" @@ -825,15 +909,10 @@ def test_chat_streaming_tool_calls_xai(mock_get_model, mock_tool_call_model): # Test Summary Function def run_all_tests(): """Run all tests and return summary""" - + # Run pytest with verbose output - exit_code = pytest.main([ - __file__, - '-v', - '--tb=short', - '--color=yes' - ]) - + exit_code = pytest.main([__file__, "-v", "--tb=short", "--color=yes"]) + return exit_code == 0 @@ -842,9 +921,9 @@ def run_all_tests(): print("TEE LLM Router - Comprehensive Unit Tests") print("=" * 80) print() - + success = run_all_tests() - + print() print("=" * 80) if success: @@ -852,5 +931,5 @@ def run_all_tests(): else: print("✗ SOME TESTS FAILED") print("=" * 80) - + sys.exit(0 if success else 1)