chart/app/api/ws.py

133 lines
5 KiB
Python

import asyncio
from fastapi import APIRouter, WebSocket, WebSocketDisconnect
from app.api.deps import token_matches
from app.bars.models import Timeframe
from app.analysis.alerts import AlertEngine
from app.analysis.confluence import cluster_levels
from app.notify.ntfy import send_ntfy
router = APIRouter()
def enabled_levels(runtime, prefs: dict | None):
if not prefs or prefs.get("hidden_levels_score"):
return runtime.levels
enabled = prefs.get("enabled", {})
ma = enabled.get("ma", {})
return [
level
for level in runtime.levels
if (level.kind.value == "ma" and level.period in ma.get(level.tf.value, []))
or (level.kind.value == "manual" and enabled.get("manual", True))
or (level.kind.value == "trendline" and enabled.get("auto", False))
]
def connection_clusters(runtime, prefs: dict | None):
if runtime.price is None or runtime.stream.last_bar_t is None:
return []
return cluster_levels(
enabled_levels(runtime, prefs), runtime.stream.last_bar_t, runtime.price, runtime.atr15
)
def snapshot(runtime, tf: Timeframe, prefs: dict | None = None) -> dict:
return {
"type": "snapshot",
"tf": tf.value,
"bars": [bar.to_dict() for bar in runtime.store.get(tf, 1000)],
"levels": [level.to_dict() for level in runtime.levels],
"clusters": [cluster.to_dict() for cluster in connection_clusters(runtime, prefs)],
"price": runtime.store.get(Timeframe.M1, 1)[-1].c
if runtime.store.get(Timeframe.M1, 1)
else None,
}
@router.websocket("/ws")
async def websocket_endpoint(websocket: WebSocket):
# Browsers cannot set headers on a WebSocket handshake, so the token comes
# in as a query parameter here. 1008 = policy violation.
if not token_matches(websocket.app, websocket.query_params.get("token", "")):
await websocket.close(code=1008, reason="Missing or invalid chart token")
return
await websocket.accept()
runtime = websocket.app.state.runtime
queue: asyncio.Queue = asyncio.Queue(maxsize=100)
runtime.subscribers.add(queue)
tf = Timeframe.M1
prefs = None
alert_engine = AlertEngine(
runtime.settings.confluence_min_score, runtime.settings.alert_cooldown_seconds
)
await websocket.send_json(snapshot(runtime, tf, prefs))
async def receive():
nonlocal tf, prefs
try:
while True:
message = await websocket.receive_json()
if message.get("type") == "subscribe":
tf = Timeframe(message.get("tf", "1m"))
await websocket.send_json(snapshot(runtime, tf, prefs))
elif message.get("type") == "prefs":
prefs = message
clusters = connection_clusters(runtime, prefs)
await websocket.send_json(
{
"type": "clusters",
"price": runtime.price,
"clusters": [cluster.to_dict() for cluster in clusters],
}
)
except WebSocketDisconnect:
queue.put_nowait({"type": "disconnect"})
receiver = asyncio.create_task(receive())
try:
while True:
event = await queue.get()
if event["type"] == "disconnect":
break
if event["type"] == "bar" and event["bar"].tf is tf:
await websocket.send_json(
{"type": "bar", "tf": tf.value, "bar": event["bar"].to_dict()}
)
elif event["type"] == "levels":
await websocket.send_json(
{"type": "levels", "levels": [level.to_dict() for level in event["levels"]]}
)
elif event["type"] == "clusters":
clusters = connection_clusters(runtime, prefs)
await websocket.send_json(
{
"type": "clusters",
"price": runtime.price,
"clusters": [cluster.to_dict() for cluster in clusters],
}
)
alerts = (
alert_engine.evaluate(
clusters,
runtime.price,
runtime.atr15,
runtime.stream.last_bar_t or 0,
runtime.stream.symbol,
)
if event.get("evaluate_alerts")
else []
)
for alert in alerts:
await websocket.send_json(
{"type": "alert", "cluster": alert.cluster.to_dict(), "message": alert.message}
)
await send_ntfy(
runtime.settings.ntfy_server, runtime.settings.ntfy_topic, alert.message
)
except (WebSocketDisconnect, asyncio.CancelledError):
pass
finally:
receiver.cancel()
runtime.subscribers.discard(queue)