Files
madturtle/python/inference_server/server.py
T

159 lines
5.5 KiB
Python
Raw Normal View History

"""
Mad Turtle Inference Server
FastAPI service that loads ONNX ensemble models and exposes REST endpoints
for MT5 EA to get BUY/SELL/HOLD signals + confidence scores.
"""
import json
import logging
from datetime import datetime, timezone
from pathlib import Path
from typing import Optional
import numpy as np
import onnxruntime as ort
from fastapi import FastAPI, HTTPException
from pydantic import BaseModel, Field
logging.basicConfig(level=logging.INFO)
logger = logging.getLogger("mad_turtle_server")
ROOT = Path(__file__).resolve().parents[2]
MODELS_DIR = ROOT / "models"
META_PATH = MODELS_DIR / "metadata.json"
app = FastAPI(title="Mad Turtle Inference Server", version="2.0.0")
class HealthResponse(BaseModel):
status: str
models_loaded: int
uptime_seconds: float
class SignalRequest(BaseModel):
features: list[float] = Field(..., min_length=14, max_length=14)
model: Optional[str] = "ensemble"
class SignalResponse(BaseModel):
signal: str
confidence: float
buy_prob: float
sell_prob: float
hold_prob: float
model_version: str
class OHLCVRequest(BaseModel):
open: float
high: float
low: float
close: float
volume: float
class Engine:
def __init__(self):
self.sessions: dict[str, ort.InferenceSession] = {}
self.metadata: dict = {}
self.feature_names: list[str] = []
self.started_at: Optional[str] = None
def load(self):
self.started_at = datetime.now(timezone.utc).isoformat()
if not META_PATH.exists():
raise FileNotFoundError(f"metadata.json not found at {META_PATH}. Run build_onnx_raw.py first.")
with open(META_PATH) as f:
self.metadata = json.load(f)
self.feature_names = self.metadata["features"]
for name, info in self.metadata.get("models", {}).items():
path = ROOT / info["path"]
if not path.exists():
logger.warning("Model file missing: %s", path)
continue
sess = ort.InferenceSession(str(path), providers=["CPUExecutionProvider"])
self.sessions[name] = sess
logger.info("Loaded model '%s' from %s", name, path)
def predict(self, features: list[float], model_name: str = "ensemble") -> dict:
if model_name not in self.sessions:
raise ValueError(f"Model '{model_name}' not loaded. Available: {list(self.sessions.keys())}")
if len(features) != len(self.feature_names):
raise ValueError(f"Expected {len(self.feature_names)} features, got {len(features)}")
x = np.array([features], dtype=np.float32)
sess = self.sessions[model_name]
input_name = sess.get_inputs()[0].name
outputs = sess.run(None, {input_name: x})[0]
probs = outputs[0]
classes = ["SELL", "HOLD", "BUY"]
idx = int(np.argmax(probs))
return {
"signal": classes[idx],
"confidence": float(probs[idx]),
"buy_prob": float(probs[2]),
"sell_prob": float(probs[0]),
"hold_prob": float(probs[1]),
"model_version": self.metadata.get("built_at", "unknown"),
}
def engineer_features(self, ohlcv: dict) -> list[float]:
import pandas as pd
df = pd.DataFrame([ohlcv])
df["returns_1"] = np.log(df["close"] / df["open"])
df["returns_3"] = np.log(df["close"] / df["close"])
df["returns_6"] = np.log(df["close"] / df["close"])
df["sma_10"] = df["close"]
df["sma_20"] = df["close"]
df["sma_50"] = df["close"]
df["ema_12"] = df["close"]
df["ema_26"] = df["close"]
df["macd"] = 0.0
df["macd_signal"] = 0.0
delta = df["close"].diff().fillna(0)
gain = delta.clip(lower=0).rolling(14).mean().fillna(0)
loss = (-delta.clip(upper=0)).rolling(14).mean().fillna(0)
rs = gain / (loss + 1e-9)
df["rsi_14"] = (100.0 - (100.0 / (1.0 + rs))).fillna(50.0)
df["atr_14"] = (df["high"] - df["low"]).fillna(0.0)
df["atr_pct"] = (df["atr_14"] / (df["close"] + 1e-9)).fillna(0.0)
df["vol_ratio"] = 1.0
df["high_low_range"] = ((df["high"] - df["low"]) / (df["close"] + 1e-9)).fillna(0.0)
df["dist_sma20"] = 0.0
row = df.iloc[-1]
return [float(row[c]) for c in self.feature_names]
engine = Engine()
@app.on_event("startup")
async def startup():
engine.load()
@app.get("/health", response_model=HealthResponse)
async def health():
now = datetime.now(timezone.utc)
start = datetime.fromisoformat(engine.started_at) if engine.started_at else now
return HealthResponse(
status="ok" if engine.sessions else "degraded",
models_loaded=len(engine.sessions),
uptime_seconds=(now - start).total_seconds(),
)
@app.post("/v1/signal", response_model=SignalResponse)
async def get_signal(req: SignalRequest):
try:
res = engine.predict(req.features, req.model or "ensemble")
return SignalResponse(**res)
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
@app.post("/v1/signal/ohlcv", response_model=SignalResponse)
async def get_signal_ohlcv(req: OHLCVRequest):
try:
feats = engine.engineer_features(req.dict())
res = engine.predict(feats)
return SignalResponse(**res)
except Exception as e:
raise HTTPException(status_code=400, detail=str(e))
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="0.0.0.0", port=8000, log_level="info")