159 lines
5.5 KiB
Python
159 lines
5.5 KiB
Python
"""
|
|
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")
|