Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension


Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
1 change: 1 addition & 0 deletions backend/app/api/__init__.py
Original file line number Diff line number Diff line change
@@ -0,0 +1 @@
"""API route modules."""
189 changes: 189 additions & 0 deletions backend/app/api/chat.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,189 @@
"""AI chat API endpoint."""
from __future__ import annotations

import json
import logging
import os

from fastapi import APIRouter, HTTPException
from pydantic import BaseModel

from app import db
from app import state

logger = logging.getLogger(__name__)
router = APIRouter(prefix="/api/chat", tags=["chat"])

_LLM_MOCK = os.environ.get("LLM_MOCK", "").lower() == "true"

_MOCK_RESPONSE = {
"message": (
"I'm analyzing your portfolio. You have a diversified mix of holdings. "
"Overall looking healthy — would you like me to suggest any rebalancing?"
),
"trades": [],
"watchlist_changes": [],
}

_SYSTEM_PROMPT = """\
You are FinAlly, an AI trading assistant for a simulated stock trading workstation.
Help users analyze their portfolio, suggest trades, and manage their watchlist.

Always respond with a single JSON object matching exactly this schema:
{
"message": "<your conversational response>",
"trades": [{"ticker": "SYMBOL", "side": "buy|sell", "quantity": <number>}],
"watchlist_changes": [{"ticker": "SYMBOL", "action": "add|remove"}]
}

Rules:
- Be concise and data-driven.
- Only execute trades when the user explicitly asks.
- Acknowledge this is a simulated portfolio with fake money.
- Respond with valid JSON only — no extra text."""


class ChatRequest(BaseModel):
message: str


async def _call_llm(messages: list[dict]) -> dict:
try:
import litellm

response = await litellm.acompletion(
model="openrouter/openai/gpt-oss-120b",
messages=messages,
api_base="https://openrouter.ai/api/v1",
api_key=os.environ.get("OPENROUTER_API_KEY", ""),
response_format={"type": "json_object"},
temperature=0.7,
extra_headers={
"X-Title": "FinAlly",
"HTTP-Referer": "https://github.com/eluiken/finally",
},
)
return json.loads(response.choices[0].message.content)
except Exception:
logger.exception("LLM call failed")
return {
"message": "I encountered an error processing your request. Please try again.",
"trades": [],
"watchlist_changes": [],
}


@router.post("")
async def chat(req: ChatRequest) -> dict:
"""Send a message; receive a response with optional auto-executed trades."""
if not req.message.strip():
raise HTTPException(status_code=422, detail="Message cannot be empty")

portfolio = await db.get_portfolio()
cash: float = portfolio["cash"]
positions: list[dict] = portfolio["positions"]
watchlist_tickers = await db.get_watchlist()

pos_lines = []
positions_value = 0.0
for pos in positions:
ticker = pos["ticker"]
qty = pos["quantity"]
avg = pos["avg_cost"]
cur = state.price_cache.get_price(ticker) or avg
val = qty * cur
pnl = (cur - avg) * qty
positions_value += val
pos_lines.append(
f" {ticker}: {qty:.2f}sh @ avg ${avg:.2f}, now ${cur:.2f}, P&L ${pnl:+.2f}"
)

wl_lines = []
for t in watchlist_tickers:
u = state.price_cache.get(t)
wl_lines.append(f"{t}:${u.price:.2f}" if u else f"{t}:N/A")

context = (
f"Portfolio: cash=${cash:.2f}, total=${cash + positions_value:.2f}\n"
f"Positions:\n" + ("\n".join(pos_lines) if pos_lines else " (none)") + "\n"
f"Watchlist: {', '.join(wl_lines)}"
)

history = await db.get_chat_history(limit=20)
messages: list[dict] = [
{"role": "system", "content": _SYSTEM_PROMPT},
{"role": "user", "content": context},
]
for msg in history:
messages.append({"role": msg["role"], "content": msg["content"]})
messages.append({"role": "user", "content": req.message})

await db.save_message("user", req.message)

llm_result = _MOCK_RESPONSE if _LLM_MOCK else await _call_llm(messages)

response_message = llm_result.get("message", "")
trades = llm_result.get("trades", [])
wl_changes = llm_result.get("watchlist_changes", [])

executed: list[dict] = []
errors: list[str] = []
for trade in trades:
ticker = str(trade.get("ticker", "")).upper()
side = str(trade.get("side", "")).lower()
try:
quantity = float(trade.get("quantity", 0))
except (TypeError, ValueError):
errors.append(f"Invalid quantity in trade: {trade}")
continue
if not ticker or side not in ("buy", "sell") or quantity <= 0:
errors.append(f"Invalid trade spec: {trade}")
continue
price = state.price_cache.get_price(ticker)
if price is None:
errors.append(f"No price available for {ticker}")
continue
result = await db.execute_trade(ticker, side, quantity, price)
if "error" in result:
errors.append(result["error"])
else:
executed.append(result)

applied_wl: list[dict] = []
for change in wl_changes:
ticker = str(change.get("ticker", "")).upper()
action = str(change.get("action", "")).lower()
if action == "add":
if await db.add_ticker(ticker) and state.market_source:
await state.market_source.add_ticker(ticker)
applied_wl.append({"ticker": ticker, "action": "add"})
elif action == "remove":
if await db.remove_ticker(ticker):
if state.market_source:
await state.market_source.remove_ticker(ticker)
state.price_cache.remove(ticker)
applied_wl.append({"ticker": ticker, "action": "remove"})

actions = None
if executed or errors or applied_wl:
actions = json.dumps({"trades": executed, "errors": errors, "watchlist_changes": applied_wl})

if executed:
try:
updated = await db.get_portfolio()
pv = sum(
p["quantity"] * (state.price_cache.get_price(p["ticker"]) or p["avg_cost"])
for p in updated["positions"]
)
await db.record_snapshot(updated["cash"] + pv)
except Exception:
logger.exception("Failed to record post-trade snapshot")

await db.save_message("assistant", response_message, actions)

return {
"message": response_message,
"trades": executed,
"trade_errors": errors,
"watchlist_changes": applied_wl,
}
9 changes: 9 additions & 0 deletions backend/app/api/health.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,9 @@
"""Health check endpoint."""
from fastapi import APIRouter

router = APIRouter(prefix="/api", tags=["system"])


@router.get("/health")
async def health() -> dict:
return {"status": "ok"}
109 changes: 109 additions & 0 deletions backend/app/api/portfolio.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,109 @@
"""Portfolio REST API endpoints."""
from __future__ import annotations

import logging
from typing import Annotated

from fastapi import APIRouter, HTTPException, Query
from pydantic import BaseModel, Field

from app import db
from app import state

logger = logging.getLogger(__name__)

router = APIRouter(prefix="/api/portfolio", tags=["portfolio"])


class TradeRequest(BaseModel):
ticker: str
side: str
quantity: float = Field(gt=0)


@router.get("")
async def get_portfolio() -> dict:
"""Return positions, cash, total value, and unrealized P&L."""
data = await db.get_portfolio()
cash: float = data["cash"]
positions: list[dict] = data["positions"]

enriched = []
positions_value = 0.0
for pos in positions:
ticker = pos["ticker"]
qty: float = pos["quantity"]
avg_cost: float = pos["avg_cost"]
current = state.price_cache.get_price(ticker)

if current is not None:
market_value = qty * current
unrealized_pnl = (current - avg_cost) * qty
pnl_percent = (current - avg_cost) / avg_cost * 100
else:
market_value = qty * avg_cost
unrealized_pnl = 0.0
pnl_percent = 0.0

positions_value += market_value
enriched.append(
{
"ticker": ticker,
"quantity": round(qty, 4),
"avg_cost": round(avg_cost, 4),
"current_price": current,
"market_value": round(market_value, 2),
"unrealized_pnl": round(unrealized_pnl, 2),
"pnl_percent": round(pnl_percent, 2),
"updated_at": pos["updated_at"],
}
)

total_value = cash + positions_value
return {
"cash": round(cash, 2),
"positions": enriched,
"positions_value": round(positions_value, 2),
"total_value": round(total_value, 2),
}


@router.post("/trade")
async def execute_trade(req: TradeRequest) -> dict:
"""Execute a market order at the current price."""
ticker = req.ticker.upper().strip()
side = req.side.lower()

if side not in ("buy", "sell"):
raise HTTPException(status_code=422, detail="side must be 'buy' or 'sell'")

price = state.price_cache.get_price(ticker)
if price is None:
raise HTTPException(status_code=422, detail=f"No price available for {ticker}")

result = await db.execute_trade(ticker, side, req.quantity, price)

if "error" in result:
raise HTTPException(status_code=422, detail=result["error"])

# Record a portfolio snapshot immediately after the trade
try:
portfolio = await db.get_portfolio()
pv = sum(
pos["quantity"] * (state.price_cache.get_price(pos["ticker"]) or pos["avg_cost"])
for pos in portfolio["positions"]
)
await db.record_snapshot(portfolio["cash"] + pv)
except Exception:
logger.exception("Failed to record snapshot after trade")

return result


@router.get("/history")
async def get_history(
limit: Annotated[int, Query(ge=1, le=5000)] = 500,
) -> dict:
"""Return portfolio value snapshots over time."""
snapshots = await db.get_snapshots(limit)
return {"snapshots": snapshots}
65 changes: 65 additions & 0 deletions backend/app/api/watchlist.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,65 @@
"""Watchlist REST API endpoints."""
from __future__ import annotations

from fastapi import APIRouter, HTTPException
from pydantic import BaseModel

from app import db
from app import state

router = APIRouter(prefix="/api/watchlist", tags=["watchlist"])


class AddTickerRequest(BaseModel):
ticker: str


@router.get("")
async def get_watchlist() -> dict:
"""Return all watchlist tickers with current prices."""
tickers = await db.get_watchlist()
result = []
for ticker in tickers:
update = state.price_cache.get(ticker)
result.append(
{
"ticker": ticker,
"price": update.price if update else None,
"change": update.change if update else None,
"change_percent": update.change_percent if update else None,
"direction": update.direction if update else None,
}
)
return {"tickers": result}


@router.post("", status_code=201)
async def add_ticker(req: AddTickerRequest) -> dict:
"""Add a ticker to the watchlist and start price tracking."""
ticker = req.ticker.upper().strip()
if not ticker.isalpha() or len(ticker) > 10:
raise HTTPException(status_code=422, detail="Invalid ticker symbol")

added = await db.add_ticker(ticker)
if not added:
raise HTTPException(status_code=409, detail=f"{ticker} already in watchlist")

if state.market_source is not None:
await state.market_source.add_ticker(ticker)

return {"ticker": ticker, "added": True}


@router.delete("/{ticker}", status_code=200)
async def remove_ticker(ticker: str) -> dict:
"""Remove a ticker from the watchlist and stop price tracking."""
ticker = ticker.upper().strip()
removed = await db.remove_ticker(ticker)
if not removed:
raise HTTPException(status_code=404, detail=f"{ticker} not in watchlist")

if state.market_source is not None:
await state.market_source.remove_ticker(ticker)
state.price_cache.remove(ticker)

return {"ticker": ticker, "removed": True}
Loading
Loading