This commit is contained in:
73 files changed
+10609
No files matched your search
+170
@@ -0,0 +1,170 @@
|
||||
package app
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/app/client"
|
||||
"github.com/pheinrich/aitrade/pkg/app/news"
|
||||
"github.com/pheinrich/aitrade/pkg/app/strategy"
|
||||
"github.com/pheinrich/aitrade/pkg/app/trader"
|
||||
"github.com/pheinrich/aitrade/pkg/app/web"
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"golang.org/x/sync/errgroup"
|
||||
)
|
||||
|
||||
type Application struct {
|
||||
cfg *config.Config
|
||||
db db.Database
|
||||
ibClient *client.IBClient
|
||||
newsAgg *news.Aggregator
|
||||
trader *trader.Trader
|
||||
webServer *web.Server
|
||||
balanceRepo *db.BalanceRepository
|
||||
newsRepo *db.NewsRepository
|
||||
tradeRepo *db.TradeRepository
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func New(cfg *config.Config, logger *slog.Logger) (*Application, error) {
|
||||
// Open database
|
||||
database, err := db.New(cfg.Database.Path, logger)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to initialize database: %w", err)
|
||||
}
|
||||
|
||||
// Create IB client
|
||||
ibClient := client.New(&cfg.IBGateway, logger)
|
||||
|
||||
// Create repositories
|
||||
balanceRepo := db.NewBalanceRepository(database)
|
||||
newsRepo := db.NewNewsRepository(database)
|
||||
tradeRepo := db.NewTradeRepository(database)
|
||||
whitelistRepo := db.NewWhitelistRepository(database)
|
||||
positionRepo := db.NewPositionRepository(database)
|
||||
|
||||
// Create LLM scorer (optional)
|
||||
var llmScorer *news.LLMScorer
|
||||
if cfg.LLMScorer.Enabled {
|
||||
llmScorer = news.NewLLMScorer(
|
||||
&cfg.LLMScorer,
|
||||
news.NewAnalyzer(), // Fallback analyzer
|
||||
logger,
|
||||
)
|
||||
logger.Info("LLM scorer enabled",
|
||||
slog.String("model", cfg.LLMScorer.ModelName),
|
||||
slog.String("endpoint", cfg.LLMScorer.Endpoint),
|
||||
slog.Float64("ensemble_weight", cfg.LLMScorer.EnsembleWeight))
|
||||
}
|
||||
|
||||
// Create news aggregator
|
||||
newsAgg := news.NewAggregator(newsRepo, cfg.News.PollInterval.Duration, llmScorer, logger)
|
||||
|
||||
// Add news sources from config
|
||||
if len(cfg.News.Sources) > 0 {
|
||||
newsAgg.AddSourcesFromConfig(cfg.News.Sources, cfg.News.DefaultRateLimit)
|
||||
}
|
||||
|
||||
// Create trading strategy
|
||||
tradingStrategy := strategy.NewStrategy(
|
||||
cfg.Trading.Strategy,
|
||||
cfg.Trading.StopLossEnabled,
|
||||
cfg.Trading.StopLossPercent,
|
||||
)
|
||||
|
||||
// Create trader
|
||||
traderInstance := trader.NewTrader(
|
||||
ibClient,
|
||||
tradeRepo,
|
||||
newsRepo,
|
||||
balanceRepo,
|
||||
whitelistRepo,
|
||||
positionRepo,
|
||||
tradingStrategy,
|
||||
&cfg.Trading,
|
||||
logger,
|
||||
)
|
||||
|
||||
// Create web server
|
||||
webServer, err := web.NewServer(
|
||||
&cfg.Web,
|
||||
&cfg.Trading,
|
||||
&cfg.OIDC,
|
||||
tradeRepo,
|
||||
balanceRepo,
|
||||
newsRepo,
|
||||
whitelistRepo,
|
||||
traderInstance,
|
||||
logger,
|
||||
)
|
||||
if err != nil {
|
||||
database.Close()
|
||||
return nil, fmt.Errorf("failed to create web server: %w", err)
|
||||
}
|
||||
|
||||
app := &Application{
|
||||
cfg: cfg,
|
||||
db: database,
|
||||
ibClient: ibClient,
|
||||
newsAgg: newsAgg,
|
||||
trader: traderInstance,
|
||||
webServer: webServer,
|
||||
balanceRepo: balanceRepo,
|
||||
newsRepo: newsRepo,
|
||||
tradeRepo: tradeRepo,
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
// Connect news aggregator to web server for SSE notifications
|
||||
newsAgg.SetNewsUpdateCallback(func() {
|
||||
webServer.BroadcastSSE("news.updated")
|
||||
})
|
||||
|
||||
return app, nil
|
||||
}
|
||||
|
||||
func (a *Application) Run(ctx context.Context) error {
|
||||
a.logger.Info("application starting")
|
||||
|
||||
g, ctx := errgroup.WithContext(ctx)
|
||||
|
||||
// Start IB Gateway client
|
||||
g.Go(func() error {
|
||||
return a.ibClient.Run(ctx)
|
||||
})
|
||||
|
||||
// Start news aggregator
|
||||
g.Go(func() error {
|
||||
return a.newsAgg.Run(ctx)
|
||||
})
|
||||
|
||||
// Start trader
|
||||
g.Go(func() error {
|
||||
return a.trader.Run(ctx)
|
||||
})
|
||||
|
||||
// Start web server
|
||||
g.Go(func() error {
|
||||
return a.webServer.Run(ctx)
|
||||
})
|
||||
|
||||
err := g.Wait()
|
||||
if err != nil && err != context.Canceled {
|
||||
return err
|
||||
}
|
||||
|
||||
a.logger.Info("shutting down")
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Application) Close() error {
|
||||
if a.ibClient != nil {
|
||||
a.ibClient.Disconnect()
|
||||
}
|
||||
if a.db != nil {
|
||||
return a.db.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,396 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
"github.com/scmhub/ibapi"
|
||||
)
|
||||
|
||||
type Client interface {
|
||||
Connect(ctx context.Context) error
|
||||
Disconnect() error
|
||||
IsConnected() bool
|
||||
GetAccountSummary(ctx context.Context) (*AccountSummary, error)
|
||||
PlaceOrder(ctx context.Context, order *Order) (int64, error)
|
||||
CancelOrder(ctx context.Context, orderID int64) error
|
||||
GetMarketData(ctx context.Context, symbol string) (*MarketData, error)
|
||||
}
|
||||
|
||||
type IBClient struct {
|
||||
cfg *config.IBGatewayConfig
|
||||
logger *slog.Logger
|
||||
|
||||
ibClient *ibapi.EClient
|
||||
connected bool
|
||||
readyChan chan struct{} // Signal when connection is truly ready
|
||||
isReady bool // True only after ConnectAck received
|
||||
isReadyMu sync.RWMutex
|
||||
|
||||
// Market data tracking
|
||||
marketDataMu sync.RWMutex
|
||||
marketData map[string]*MarketData // symbol -> latest market data
|
||||
reqIDToSymbol map[int64]string // reqID -> symbol mapping
|
||||
reqIDGen *RequestIDGenerator // Sequential request ID generator
|
||||
|
||||
// Account summary tracking
|
||||
accountSummaryMu sync.RWMutex
|
||||
accountSummary *AccountSummary
|
||||
accountSummaryReqID int64 // Subscription request ID
|
||||
accountSummaryActive bool // Is subscription active
|
||||
}
|
||||
|
||||
type AccountSummary struct {
|
||||
TotalValue float64
|
||||
CashBalance float64
|
||||
BuyingPower float64
|
||||
UnrealizedPnL float64
|
||||
RealizedPnL float64
|
||||
}
|
||||
|
||||
type Order struct {
|
||||
Symbol string
|
||||
Action model.ActionType
|
||||
Quantity int
|
||||
OrderType string // MKT, LMT, STP
|
||||
LimitPrice *float64
|
||||
StopPrice *float64
|
||||
}
|
||||
|
||||
type MarketData struct {
|
||||
Symbol string
|
||||
LastPrice float64
|
||||
BidPrice float64
|
||||
AskPrice float64
|
||||
Volume int64
|
||||
Timestamp time.Time
|
||||
}
|
||||
|
||||
func New(cfg *config.IBGatewayConfig, logger *slog.Logger) *IBClient {
|
||||
return &IBClient{
|
||||
cfg: cfg,
|
||||
logger: logger,
|
||||
connected: false,
|
||||
readyChan: make(chan struct{}),
|
||||
marketData: make(map[string]*MarketData),
|
||||
reqIDToSymbol: make(map[int64]string),
|
||||
reqIDGen: NewRequestIDGenerator(logger),
|
||||
accountSummary: &AccountSummary{},
|
||||
}
|
||||
}
|
||||
|
||||
func (c *IBClient) Connect(ctx context.Context) error {
|
||||
c.logger.Info("connecting to IB Gateway",
|
||||
slog.String("host", c.cfg.Host),
|
||||
slog.Int("port", c.cfg.Port),
|
||||
slog.Int("client_id", c.cfg.ClientID),
|
||||
)
|
||||
|
||||
// Create IB client with custom wrapper for market data callbacks
|
||||
wrapper := NewCustomWrapper(c.logger, c)
|
||||
ibClient := ibapi.NewEClient(wrapper)
|
||||
|
||||
// Connect with retry logic
|
||||
maxRetries := 5
|
||||
backoff := time.Second
|
||||
|
||||
for attempt := 1; attempt <= maxRetries; attempt++ {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
default:
|
||||
}
|
||||
|
||||
err := ibClient.Connect(c.cfg.Host, c.cfg.Port, int64(c.cfg.ClientID))
|
||||
if err == nil {
|
||||
c.ibClient = ibClient
|
||||
c.connected = true
|
||||
c.logger.Info("connected to IB Gateway")
|
||||
|
||||
// Wait for ConnectAck from IB Gateway (or timeout after 2 seconds)
|
||||
// scmhub/ibapi runs message loop automatically - no need to call Run()
|
||||
select {
|
||||
case <-c.readyChan:
|
||||
c.isReadyMu.Lock()
|
||||
c.isReady = true
|
||||
c.isReadyMu.Unlock()
|
||||
c.logger.Info("IB Gateway ready for market data requests")
|
||||
case <-time.After(2 * time.Second):
|
||||
c.logger.Warn("ConnectAck timeout - proceeding anyway")
|
||||
c.isReadyMu.Lock()
|
||||
c.isReady = true // Allow proceeding even without ConnectAck
|
||||
c.isReadyMu.Unlock()
|
||||
}
|
||||
|
||||
// Subscribe to account summary updates (streaming)
|
||||
c.subscribeAccountSummary()
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
c.logger.Warn("connection attempt failed",
|
||||
slog.Int("attempt", attempt),
|
||||
slog.Int("max_retries", maxRetries),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
|
||||
if attempt < maxRetries {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-time.After(backoff):
|
||||
backoff = min(backoff*2, 60*time.Second)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return fmt.Errorf("failed to connect after %d attempts", maxRetries)
|
||||
}
|
||||
|
||||
func (c *IBClient) Disconnect() error {
|
||||
if c.ibClient != nil && c.connected {
|
||||
c.logger.Info("disconnecting from IB Gateway")
|
||||
|
||||
// Cancel account summary subscription if active
|
||||
c.accountSummaryMu.Lock()
|
||||
if c.accountSummaryActive {
|
||||
c.ibClient.CancelAccountSummary(c.accountSummaryReqID)
|
||||
c.accountSummaryActive = false
|
||||
}
|
||||
c.accountSummaryMu.Unlock()
|
||||
|
||||
c.ibClient.Disconnect()
|
||||
c.connected = false
|
||||
c.isReadyMu.Lock()
|
||||
c.isReady = false
|
||||
c.isReadyMu.Unlock()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *IBClient) IsConnected() bool {
|
||||
c.isReadyMu.RLock()
|
||||
defer c.isReadyMu.RUnlock()
|
||||
return c.isReady && c.ibClient != nil
|
||||
}
|
||||
|
||||
// subscribeAccountSummary subscribes to account summary updates (streaming)
|
||||
func (c *IBClient) subscribeAccountSummary() {
|
||||
c.accountSummaryMu.Lock()
|
||||
defer c.accountSummaryMu.Unlock()
|
||||
|
||||
if c.accountSummaryActive {
|
||||
return // Already subscribed
|
||||
}
|
||||
|
||||
reqID := c.reqIDGen.Next()
|
||||
c.accountSummaryReqID = reqID
|
||||
c.accountSummaryActive = true
|
||||
|
||||
c.logger.Info("subscribing to account summary updates", slog.Int64("req_id", reqID))
|
||||
c.ibClient.ReqAccountSummary(reqID, "All", "NetLiquidation,TotalCashValue,BuyingPower,UnrealizedPnL,RealizedPnL")
|
||||
}
|
||||
|
||||
func (c *IBClient) GetAccountSummary(ctx context.Context) (*AccountSummary, error) {
|
||||
if !c.IsConnected() {
|
||||
return nil, fmt.Errorf("not connected to IB Gateway")
|
||||
}
|
||||
|
||||
// Just return the current (streamed) account summary
|
||||
c.accountSummaryMu.RLock()
|
||||
defer c.accountSummaryMu.RUnlock()
|
||||
|
||||
if c.accountSummary.TotalValue == 0 {
|
||||
return nil, fmt.Errorf("account summary not yet available")
|
||||
}
|
||||
|
||||
return c.accountSummary, nil
|
||||
}
|
||||
|
||||
func (c *IBClient) PlaceOrder(ctx context.Context, order *Order) (int64, error) {
|
||||
if !c.IsConnected() {
|
||||
return 0, fmt.Errorf("not connected to IB Gateway")
|
||||
}
|
||||
|
||||
// Generate unique order ID
|
||||
orderID := time.Now().Unix()
|
||||
|
||||
// Create IB contract
|
||||
contract := &ibapi.Contract{
|
||||
Symbol: order.Symbol,
|
||||
SecType: "STK",
|
||||
Exchange: "SMART",
|
||||
Currency: "USD",
|
||||
}
|
||||
|
||||
// Create IB order
|
||||
ibOrder := &ibapi.Order{
|
||||
Action: string(order.Action),
|
||||
TotalQuantity: ibapi.StringToDecimal(fmt.Sprintf("%d", order.Quantity)),
|
||||
OrderType: order.OrderType,
|
||||
TIF: "DAY",
|
||||
}
|
||||
|
||||
if order.LimitPrice != nil {
|
||||
ibOrder.LmtPrice = *order.LimitPrice
|
||||
}
|
||||
|
||||
if order.StopPrice != nil {
|
||||
ibOrder.AuxPrice = *order.StopPrice
|
||||
}
|
||||
|
||||
// Place order
|
||||
c.ibClient.PlaceOrder(orderID, contract, ibOrder)
|
||||
|
||||
c.logger.Info("placed order",
|
||||
slog.Int64("order_id", orderID),
|
||||
slog.String("symbol", order.Symbol),
|
||||
slog.String("action", string(order.Action)),
|
||||
slog.Int("quantity", order.Quantity),
|
||||
)
|
||||
|
||||
return orderID, nil
|
||||
}
|
||||
|
||||
func (c *IBClient) CancelOrder(ctx context.Context, orderID int64) error {
|
||||
if !c.IsConnected() {
|
||||
return fmt.Errorf("not connected to IB Gateway")
|
||||
}
|
||||
|
||||
c.ibClient.CancelOrder(orderID, ibapi.OrderCancel{})
|
||||
|
||||
c.logger.Info("cancelled order", slog.Int64("order_id", orderID))
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (c *IBClient) Run(ctx context.Context) error {
|
||||
// Connect
|
||||
if err := c.Connect(ctx); err != nil {
|
||||
return fmt.Errorf("failed to connect: %w", err)
|
||||
}
|
||||
|
||||
// Keep connection alive and handle reconnects
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
c.Disconnect()
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
if !c.IsConnected() {
|
||||
c.logger.Warn("connection lost, attempting to reconnect")
|
||||
if err := c.Connect(ctx); err != nil {
|
||||
c.logger.Error("reconnection failed", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func getMarketDataTypeName(dataType int) string {
|
||||
switch dataType {
|
||||
case 1:
|
||||
return "Live"
|
||||
case 2:
|
||||
return "Frozen"
|
||||
case 3:
|
||||
return "Delayed (15min)"
|
||||
case 4:
|
||||
return "Delayed-Frozen"
|
||||
default:
|
||||
return fmt.Sprintf("Unknown (%d)", dataType)
|
||||
}
|
||||
}
|
||||
|
||||
// GetMarketData requests real-time market data for a symbol from IB Gateway
|
||||
func (c *IBClient) GetMarketData(ctx context.Context, symbol string) (*MarketData, error) {
|
||||
if !c.IsConnected() {
|
||||
return nil, fmt.Errorf("not connected to IB Gateway")
|
||||
}
|
||||
|
||||
// Check if we have recent cached data (within last 5 seconds)
|
||||
c.marketDataMu.RLock()
|
||||
if data, ok := c.marketData[symbol]; ok {
|
||||
if time.Since(data.Timestamp) < 5*time.Second {
|
||||
c.marketDataMu.RUnlock()
|
||||
return data, nil
|
||||
}
|
||||
}
|
||||
c.marketDataMu.RUnlock()
|
||||
|
||||
// Create contract for the symbol (US stocks on SMART exchange)
|
||||
contract := &ibapi.Contract{
|
||||
Symbol: symbol,
|
||||
SecType: "STK",
|
||||
Exchange: "SMART",
|
||||
Currency: "USD",
|
||||
}
|
||||
|
||||
// Request market data type from config
|
||||
// 1=Live, 2=Frozen, 3=Delayed(15min), 4=Delayed-Frozen
|
||||
c.ibClient.ReqMarketDataType(int64(c.cfg.MarketDataType))
|
||||
c.logger.Info("requesting market data",
|
||||
slog.String("type", getMarketDataTypeName(c.cfg.MarketDataType)))
|
||||
|
||||
// Generate sequential request ID
|
||||
reqID := c.reqIDGen.Next()
|
||||
|
||||
c.marketDataMu.Lock()
|
||||
c.reqIDToSymbol[reqID] = symbol
|
||||
c.marketDataMu.Unlock()
|
||||
|
||||
c.logger.Info("requesting market data snapshot",
|
||||
slog.String("symbol", symbol),
|
||||
slog.String("exchange", "SMART"),
|
||||
slog.Int64("req_id", reqID))
|
||||
|
||||
// genericTickList: empty string means all available ticks
|
||||
// snapshot = FALSE: streaming data (old TWS may not support snapshots)
|
||||
// regulatorySnapshot = false: not regulatory snapshot
|
||||
c.ibClient.ReqMktData(reqID, contract, "", false, false, nil)
|
||||
|
||||
c.logger.Info("market data request sent (streaming mode, live data)")
|
||||
|
||||
// Wait for data with timeout
|
||||
timeout := time.After(5 * time.Second) // Increased timeout for streaming
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
// Track if we should cancel the market data subscription
|
||||
var dataReceived bool
|
||||
defer func() {
|
||||
if dataReceived {
|
||||
// Cancel streaming market data subscription
|
||||
c.ibClient.CancelMktData(reqID)
|
||||
c.logger.Debug("cancelled market data subscription", slog.Int64("req_id", reqID))
|
||||
}
|
||||
}()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
return nil, ctx.Err()
|
||||
case <-timeout:
|
||||
return nil, fmt.Errorf("timeout waiting for market data for %s", symbol)
|
||||
case <-ticker.C:
|
||||
// Check if data arrived
|
||||
c.marketDataMu.RLock()
|
||||
if data, ok := c.marketData[symbol]; ok {
|
||||
if time.Since(data.Timestamp) < 1*time.Second {
|
||||
c.marketDataMu.RUnlock()
|
||||
dataReceived = true
|
||||
return data, nil
|
||||
}
|
||||
}
|
||||
c.marketDataMu.RUnlock()
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,54 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"sync"
|
||||
)
|
||||
|
||||
// RequestIDGenerator generates sequential request IDs for IB API calls
|
||||
type RequestIDGenerator struct {
|
||||
mu sync.Mutex
|
||||
current int64
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// NewRequestIDGenerator creates a new request ID generator starting at 1
|
||||
func NewRequestIDGenerator(logger *slog.Logger) *RequestIDGenerator {
|
||||
return &RequestIDGenerator{
|
||||
current: 1,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
// Next returns the next sequential request ID and logs it
|
||||
func (r *RequestIDGenerator) Next() int64 {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
id := r.current
|
||||
r.current++
|
||||
|
||||
if r.logger != nil {
|
||||
r.logger.Debug("generated request ID", slog.Int64("req_id", id))
|
||||
}
|
||||
|
||||
return id
|
||||
}
|
||||
|
||||
// Current returns the current request ID without incrementing
|
||||
func (r *RequestIDGenerator) Current() int64 {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
return r.current
|
||||
}
|
||||
|
||||
// Reset resets the generator back to 1 (useful for testing or reconnection)
|
||||
func (r *RequestIDGenerator) Reset() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
r.current = 1
|
||||
|
||||
if r.logger != nil {
|
||||
r.logger.Debug("reset request ID generator")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,218 @@
|
||||
package client
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/scmhub/ibapi"
|
||||
)
|
||||
|
||||
// CustomWrapper wraps the default ibapi.Wrapper and overrides only the callbacks we need
|
||||
type CustomWrapper struct {
|
||||
ibapi.Wrapper // Embed default wrapper for all other methods
|
||||
logger *slog.Logger
|
||||
client *IBClient
|
||||
}
|
||||
|
||||
// NewCustomWrapper creates a new wrapper with default implementations for all EWrapper methods
|
||||
func NewCustomWrapper(logger *slog.Logger, client *IBClient) *CustomWrapper {
|
||||
return &CustomWrapper{
|
||||
Wrapper: ibapi.Wrapper{},
|
||||
logger: logger,
|
||||
client: client,
|
||||
}
|
||||
}
|
||||
|
||||
// ConnectAck - override to signal connection is ready
|
||||
func (w *CustomWrapper) ConnectAck() {
|
||||
w.logger.Info("IB connection acknowledged")
|
||||
|
||||
// Signal that connection is ready
|
||||
if w.client != nil {
|
||||
select {
|
||||
case w.client.readyChan <- struct{}{}:
|
||||
// Signal sent
|
||||
default:
|
||||
// Channel already closed or full, ignore
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// NextValidID - override for logging
|
||||
func (w *CustomWrapper) NextValidID(orderID int64) {
|
||||
w.logger.Debug("NextValidID received", slog.Int64("order_id", orderID))
|
||||
}
|
||||
|
||||
// ManagedAccounts - override for logging (scmhub uses []string, not string)
|
||||
func (w *CustomWrapper) ManagedAccounts(accountsList []string) {
|
||||
w.logger.Info("ManagedAccounts received", slog.Any("accounts", accountsList))
|
||||
}
|
||||
|
||||
// Error - override to handle IB API errors
|
||||
func (w *CustomWrapper) Error(reqID int64, errTime int64, errCode int64, errString string, advancedOrderRejectJson string) {
|
||||
w.logger.Error("IB API error",
|
||||
slog.Int64("req_id", reqID),
|
||||
slog.Int64("err_time", errTime),
|
||||
slog.Int64("code", errCode),
|
||||
slog.String("message", errString),
|
||||
)
|
||||
|
||||
// Special handling for market data errors
|
||||
if errCode == 354 {
|
||||
w.logger.Warn("Market data permission error - check your IB account subscriptions",
|
||||
slog.Int64("req_id", reqID))
|
||||
}
|
||||
}
|
||||
|
||||
// ConnectionClosed - override for logging
|
||||
func (w *CustomWrapper) ConnectionClosed() {
|
||||
w.logger.Warn("IB connection closed")
|
||||
}
|
||||
|
||||
// TickPrice - handle price ticks for market data
|
||||
func (w *CustomWrapper) TickPrice(reqID int64, tickType ibapi.TickType, price float64, attrib ibapi.TickAttrib) {
|
||||
w.logger.Debug("received tick price",
|
||||
slog.Int64("req_id", reqID),
|
||||
slog.Int64("tick_type", int64(tickType)),
|
||||
slog.Float64("price", price))
|
||||
|
||||
if w.client == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// Map tick types to our market data structure
|
||||
// tickType: 1=bid, 2=ask, 4=last, 6=high, 7=low, 9=close
|
||||
symbol := w.getSymbolForReqID(reqID)
|
||||
if symbol == "" {
|
||||
w.logger.Warn("received tick for unknown reqID", slog.Int64("req_id", reqID))
|
||||
return
|
||||
}
|
||||
|
||||
w.client.marketDataMu.Lock()
|
||||
defer w.client.marketDataMu.Unlock()
|
||||
|
||||
// Get or create market data entry
|
||||
data, exists := w.client.marketData[symbol]
|
||||
if !exists {
|
||||
data = &MarketData{
|
||||
Symbol: symbol,
|
||||
Timestamp: time.Now(),
|
||||
}
|
||||
w.client.marketData[symbol] = data
|
||||
}
|
||||
|
||||
// Update timestamp and data based on tick type
|
||||
data.Timestamp = time.Now()
|
||||
|
||||
switch tickType {
|
||||
case 1, 66: // Bid (Live or Delayed)
|
||||
data.BidPrice = price
|
||||
case 2, 67: // Ask (Live or Delayed)
|
||||
data.AskPrice = price
|
||||
case 4, 68: // Last (Live or Delayed)
|
||||
data.LastPrice = price
|
||||
}
|
||||
|
||||
w.logger.Info("updated market data",
|
||||
slog.String("symbol", symbol),
|
||||
slog.Float64("last", data.LastPrice),
|
||||
slog.Float64("bid", data.BidPrice),
|
||||
slog.Float64("ask", data.AskPrice))
|
||||
}
|
||||
|
||||
// TickSize - handle size ticks for market data
|
||||
func (w *CustomWrapper) TickSize(reqID int64, tickType ibapi.TickType, size ibapi.Decimal) {
|
||||
if w.client == nil {
|
||||
return
|
||||
}
|
||||
|
||||
// tickType: 0=bid size, 3=ask size, 5=last size, 8=volume
|
||||
symbol := w.getSymbolForReqID(reqID)
|
||||
if symbol == "" {
|
||||
return
|
||||
}
|
||||
|
||||
w.client.marketDataMu.Lock()
|
||||
defer w.client.marketDataMu.Unlock()
|
||||
|
||||
data, exists := w.client.marketData[symbol]
|
||||
if !exists {
|
||||
data = &MarketData{
|
||||
Symbol: symbol,
|
||||
Timestamp: time.Now(),
|
||||
}
|
||||
w.client.marketData[symbol] = data
|
||||
}
|
||||
|
||||
// Update timestamp and volume
|
||||
data.Timestamp = time.Now()
|
||||
|
||||
if tickType == 8 {
|
||||
data.Volume = size.Int()
|
||||
}
|
||||
}
|
||||
|
||||
// Helper function to extract symbol from reqID
|
||||
func (w *CustomWrapper) getSymbolForReqID(reqID int64) string {
|
||||
if w.client == nil {
|
||||
return ""
|
||||
}
|
||||
|
||||
w.client.marketDataMu.RLock()
|
||||
defer w.client.marketDataMu.RUnlock()
|
||||
|
||||
symbol, ok := w.client.reqIDToSymbol[reqID]
|
||||
if !ok {
|
||||
return ""
|
||||
}
|
||||
|
||||
return symbol
|
||||
}
|
||||
|
||||
// AccountSummary - handle account summary updates
|
||||
func (w *CustomWrapper) AccountSummary(reqID int64, account string, tag string, value string, currency string) {
|
||||
if w.client == nil {
|
||||
return
|
||||
}
|
||||
|
||||
w.logger.Debug("received account summary",
|
||||
slog.Int64("req_id", reqID),
|
||||
slog.String("account", account),
|
||||
slog.String("tag", tag),
|
||||
slog.String("value", value),
|
||||
slog.String("currency", currency))
|
||||
|
||||
w.client.accountSummaryMu.Lock()
|
||||
defer w.client.accountSummaryMu.Unlock()
|
||||
|
||||
// Parse value as float
|
||||
var floatValue float64
|
||||
if _, err := fmt.Sscanf(value, "%f", &floatValue); err != nil {
|
||||
w.logger.Warn("failed to parse account summary value", slog.String("tag", tag), slog.String("value", value))
|
||||
return
|
||||
}
|
||||
|
||||
// Update account summary based on tag
|
||||
switch tag {
|
||||
case "NetLiquidation":
|
||||
w.client.accountSummary.TotalValue = floatValue
|
||||
case "TotalCashValue":
|
||||
w.client.accountSummary.CashBalance = floatValue
|
||||
case "BuyingPower":
|
||||
w.client.accountSummary.BuyingPower = floatValue
|
||||
case "UnrealizedPnL":
|
||||
w.client.accountSummary.UnrealizedPnL = floatValue
|
||||
case "RealizedPnL":
|
||||
w.client.accountSummary.RealizedPnL = floatValue
|
||||
}
|
||||
|
||||
w.logger.Debug("updated account summary",
|
||||
slog.String("tag", tag),
|
||||
slog.Float64("value", floatValue))
|
||||
}
|
||||
|
||||
// AccountSummaryEnd - handle end of account summary
|
||||
func (w *CustomWrapper) AccountSummaryEnd(reqID int64) {
|
||||
w.logger.Debug("account summary end", slog.Int64("req_id", reqID))
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
package news
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
)
|
||||
|
||||
type Aggregator struct {
|
||||
sources []Source
|
||||
newsRepo *db.NewsRepository
|
||||
analyzer *Analyzer
|
||||
llmScorer *LLMScorer
|
||||
pollInterval time.Duration
|
||||
logger *slog.Logger
|
||||
onNewsUpdated func() // Callback for SSE notifications
|
||||
}
|
||||
|
||||
func NewAggregator(
|
||||
newsRepo *db.NewsRepository,
|
||||
pollInterval time.Duration,
|
||||
llmScorer *LLMScorer,
|
||||
logger *slog.Logger,
|
||||
) *Aggregator {
|
||||
agg := &Aggregator{
|
||||
sources: make([]Source, 0),
|
||||
newsRepo: newsRepo,
|
||||
analyzer: NewAnalyzer(),
|
||||
llmScorer: llmScorer,
|
||||
pollInterval: pollInterval,
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
return agg
|
||||
}
|
||||
|
||||
// SetNewsUpdateCallback sets a callback that's called when news are updated
|
||||
func (a *Aggregator) SetNewsUpdateCallback(callback func()) {
|
||||
a.onNewsUpdated = callback
|
||||
}
|
||||
|
||||
// AddSourcesFromConfig adds sources from config
|
||||
func (a *Aggregator) AddSourcesFromConfig(sources []config.NewsSource, defaultRateLimit *config.RateLimit) {
|
||||
for _, src := range sources {
|
||||
if !src.Enabled || src.Name == "" || src.URL == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
var source Source
|
||||
|
||||
switch src.Type {
|
||||
case "rss":
|
||||
rssSource := NewRSSSource(src.Name, src.URL, a.logger)
|
||||
|
||||
// Add auth if provided
|
||||
if src.Auth != nil && src.Auth.Type == "basic" {
|
||||
rssSource.SetBasicAuth(src.Auth.Username, src.Auth.Password)
|
||||
}
|
||||
|
||||
// Add custom headers if provided
|
||||
for key, value := range src.Headers {
|
||||
rssSource.AddHeader(key, value)
|
||||
}
|
||||
|
||||
source = rssSource
|
||||
|
||||
case "alphavantage":
|
||||
source = NewAlphaVantageSource(src.Name, src.URL, a.logger)
|
||||
|
||||
case "finnhub":
|
||||
source = NewFinnhubSource(src.Name, src.URL, a.logger)
|
||||
|
||||
default:
|
||||
a.logger.Warn("unsupported source type", slog.String("type", src.Type), slog.String("name", src.Name))
|
||||
continue
|
||||
}
|
||||
|
||||
// Set rate limits: source-specific OR default
|
||||
rateLimit := src.RateLimit
|
||||
if rateLimit == nil {
|
||||
rateLimit = defaultRateLimit
|
||||
}
|
||||
|
||||
if rateLimit != nil {
|
||||
// All source types implement SetRateLimit via their embedded rate limiter
|
||||
switch s := source.(type) {
|
||||
case *RSSSource:
|
||||
s.SetRateLimit(rateLimit.MaxPerHour, rateLimit.MaxPerDay)
|
||||
case *AlphaVantageSource:
|
||||
s.SetRateLimit(rateLimit.MaxPerHour, rateLimit.MaxPerDay)
|
||||
case *FinnhubSource:
|
||||
s.SetRateLimit(rateLimit.MaxPerHour, rateLimit.MaxPerDay)
|
||||
}
|
||||
|
||||
a.logger.Info("rate limit configured",
|
||||
slog.String("source", src.Name),
|
||||
slog.Int("max_per_hour", rateLimit.MaxPerHour),
|
||||
slog.Int("max_per_day", rateLimit.MaxPerDay))
|
||||
}
|
||||
|
||||
a.AddSource(source)
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Aggregator) AddSource(source Source) {
|
||||
a.sources = append(a.sources, source)
|
||||
a.logger.Info("added news source", slog.String("source", source.Name()))
|
||||
}
|
||||
|
||||
func (a *Aggregator) Run(ctx context.Context) error {
|
||||
a.logger.Info("news aggregator starting",
|
||||
slog.Int("sources", len(a.sources)),
|
||||
slog.Duration("poll_interval", a.pollInterval),
|
||||
)
|
||||
|
||||
// Fetch immediately on start
|
||||
if err := a.fetchAllSources(ctx); err != nil {
|
||||
a.logger.Error("initial news fetch failed", slog.Any("error", err))
|
||||
}
|
||||
|
||||
ticker := time.NewTicker(a.pollInterval)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
a.logger.Info("news aggregator stopping")
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
if err := a.fetchAllSources(ctx); err != nil {
|
||||
a.logger.Error("news fetch failed", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Aggregator) fetchAllSources(ctx context.Context) error {
|
||||
a.logger.Debug("fetching from all news sources")
|
||||
|
||||
totalFetched := 0
|
||||
totalStored := 0
|
||||
|
||||
for _, source := range a.sources {
|
||||
articles, err := source.Fetch(ctx)
|
||||
if err != nil {
|
||||
a.logger.Warn("failed to fetch from source",
|
||||
slog.String("source", source.Name()),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
totalFetched += len(articles)
|
||||
|
||||
// Analyze sentiment and store
|
||||
for _, article := range articles {
|
||||
// Always use keyword analyzer first (fast)
|
||||
a.analyzer.Analyze(article)
|
||||
|
||||
// If LLM scorer enabled, use it (may fallback to keyword)
|
||||
if a.llmScorer != nil && a.llmScorer.Enabled() {
|
||||
if err := a.llmScorer.Analyze(ctx, article); err != nil {
|
||||
a.logger.Warn("LLM scoring failed",
|
||||
slog.String("url", article.URL),
|
||||
slog.Any("error", err))
|
||||
// Article already has keyword sentiment, continue
|
||||
}
|
||||
} else {
|
||||
// No LLM scorer, mark as keyword-only
|
||||
article.SentimentMethod = "keyword"
|
||||
}
|
||||
|
||||
if err := a.newsRepo.Create(ctx, article); err != nil {
|
||||
a.logger.Error("failed to store article",
|
||||
slog.String("url", article.URL),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
totalStored++
|
||||
}
|
||||
}
|
||||
|
||||
a.logger.Info("news fetch completed",
|
||||
slog.Int("fetched", totalFetched),
|
||||
slog.Int("stored", totalStored),
|
||||
slog.Int("duplicates", totalFetched-totalStored),
|
||||
)
|
||||
|
||||
// Notify listeners if new articles were stored
|
||||
if totalStored > 0 && a.onNewsUpdated != nil {
|
||||
a.onNewsUpdated()
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (a *Aggregator) GetRecent(ctx context.Context, limit int) ([]*db.NewsRepository, error) {
|
||||
return nil, fmt.Errorf("not implemented")
|
||||
}
|
||||
@@ -0,0 +1,94 @@
|
||||
package news
|
||||
|
||||
import (
|
||||
"strings"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
// Simple keyword-based sentiment analyzer
|
||||
type Analyzer struct {
|
||||
positiveKeywords map[string]int
|
||||
negativeKeywords map[string]int
|
||||
}
|
||||
|
||||
func NewAnalyzer() *Analyzer {
|
||||
return &Analyzer{
|
||||
positiveKeywords: map[string]int{
|
||||
"profit": 2,
|
||||
"gain": 2,
|
||||
"growth": 2,
|
||||
"surge": 2,
|
||||
"rally": 2,
|
||||
"bullish": 3,
|
||||
"upgrade": 2,
|
||||
"beat": 2,
|
||||
"record": 1,
|
||||
"strong": 1,
|
||||
"positive": 1,
|
||||
"success": 1,
|
||||
"win": 1,
|
||||
"jump": 2,
|
||||
"soar": 3,
|
||||
"outperform": 2,
|
||||
},
|
||||
negativeKeywords: map[string]int{
|
||||
"loss": 2,
|
||||
"decline": 2,
|
||||
"fall": 2,
|
||||
"drop": 2,
|
||||
"crash": 3,
|
||||
"bearish": 3,
|
||||
"downgrade": 2,
|
||||
"miss": 2,
|
||||
"weak": 1,
|
||||
"negative": 1,
|
||||
"fail": 2,
|
||||
"plunge": 3,
|
||||
"slump": 2,
|
||||
"underperform": 2,
|
||||
"risk": 1,
|
||||
"concern": 1,
|
||||
},
|
||||
}
|
||||
}
|
||||
|
||||
func (a *Analyzer) Analyze(article *model.NewsArticle) {
|
||||
text := strings.ToLower(article.Title + " " + article.Content)
|
||||
words := strings.Fields(text)
|
||||
|
||||
positiveScore := 0
|
||||
negativeScore := 0
|
||||
|
||||
for _, word := range words {
|
||||
word = strings.Trim(word, ".,!?;:\"'()")
|
||||
|
||||
if score, ok := a.positiveKeywords[word]; ok {
|
||||
positiveScore += score
|
||||
}
|
||||
if score, ok := a.negativeKeywords[word]; ok {
|
||||
negativeScore += score
|
||||
}
|
||||
}
|
||||
|
||||
// Calculate sentiment score from -1.0 to 1.0
|
||||
totalScore := positiveScore + negativeScore
|
||||
if totalScore == 0 {
|
||||
score := 0.0
|
||||
article.SentimentScore = &score
|
||||
article.SentimentLabel = "neutral"
|
||||
return
|
||||
}
|
||||
|
||||
sentimentScore := float64(positiveScore-negativeScore) / float64(totalScore)
|
||||
article.SentimentScore = &sentimentScore
|
||||
|
||||
// Label the sentiment
|
||||
if sentimentScore > 0.3 {
|
||||
article.SentimentLabel = "positive"
|
||||
} else if sentimentScore < -0.3 {
|
||||
article.SentimentLabel = "negative"
|
||||
} else {
|
||||
article.SentimentLabel = "neutral"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,294 @@
|
||||
package news
|
||||
|
||||
import (
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
// AlphaVantageSource fetches news from Alpha Vantage News API
|
||||
type AlphaVantageSource struct {
|
||||
name string
|
||||
url string
|
||||
logger *slog.Logger
|
||||
httpClient *http.Client
|
||||
|
||||
// Rate limiting
|
||||
maxPerHour int
|
||||
maxPerDay int
|
||||
hourlyCounter int
|
||||
dailyCounter int
|
||||
lastHourReset time.Time
|
||||
lastDayReset time.Time
|
||||
}
|
||||
|
||||
type alphaVantageResponse struct {
|
||||
Feed []struct {
|
||||
Title string `json:"title"`
|
||||
URL string `json:"url"`
|
||||
TimePublished string `json:"time_published"`
|
||||
Summary string `json:"summary"`
|
||||
Source string `json:"source"`
|
||||
} `json:"feed"`
|
||||
}
|
||||
|
||||
func NewAlphaVantageSource(name, url string, logger *slog.Logger) *AlphaVantageSource {
|
||||
now := time.Now()
|
||||
return &AlphaVantageSource{
|
||||
name: name,
|
||||
url: url,
|
||||
logger: logger,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
lastHourReset: now,
|
||||
lastDayReset: now,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AlphaVantageSource) SetRateLimit(maxPerHour, maxPerDay int) {
|
||||
s.maxPerHour = maxPerHour
|
||||
s.maxPerDay = maxPerDay
|
||||
}
|
||||
|
||||
func (s *AlphaVantageSource) Name() string {
|
||||
return s.name
|
||||
}
|
||||
|
||||
func (s *AlphaVantageSource) Fetch(ctx context.Context) ([]*model.NewsArticle, error) {
|
||||
// Check rate limits
|
||||
if !s.checkRateLimit() {
|
||||
s.logger.Warn("rate limit reached, skipping fetch",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("hourly", s.hourlyCounter),
|
||||
slog.Int("max_per_hour", s.maxPerHour),
|
||||
slog.Int("daily", s.dailyCounter),
|
||||
slog.Int("max_per_day", s.maxPerDay))
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
s.logger.Debug("fetching Alpha Vantage news", slog.String("source", s.name))
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", s.url, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("http error: %s", resp.Status)
|
||||
}
|
||||
|
||||
var data alphaVantageResponse
|
||||
if err := json.NewDecoder(resp.Body).Decode(&data); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse JSON: %w", err)
|
||||
}
|
||||
|
||||
var articles []*model.NewsArticle
|
||||
for _, item := range data.Feed {
|
||||
// Parse time: "20260702T133830"
|
||||
publishedAt := time.Now()
|
||||
if t, err := time.Parse("20060102T150405", item.TimePublished); err == nil {
|
||||
publishedAt = t
|
||||
}
|
||||
|
||||
article := &model.NewsArticle{
|
||||
Source: s.name,
|
||||
Title: item.Title,
|
||||
URL: item.URL,
|
||||
Content: item.Summary,
|
||||
PublishedAt: publishedAt,
|
||||
FetchedAt: time.Now(),
|
||||
Symbols: extractSymbols(item.Title + " " + item.Summary),
|
||||
}
|
||||
|
||||
articles = append(articles, article)
|
||||
}
|
||||
|
||||
s.logger.Debug("fetched articles",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("count", len(articles)))
|
||||
|
||||
s.incrementCounters()
|
||||
return articles, nil
|
||||
}
|
||||
|
||||
func (s *AlphaVantageSource) checkRateLimit() bool {
|
||||
now := time.Now()
|
||||
|
||||
if now.Sub(s.lastHourReset) >= time.Hour {
|
||||
s.hourlyCounter = 0
|
||||
s.lastHourReset = now
|
||||
}
|
||||
|
||||
if now.Sub(s.lastDayReset) >= 24*time.Hour {
|
||||
s.dailyCounter = 0
|
||||
s.lastDayReset = now
|
||||
}
|
||||
|
||||
if s.maxPerHour > 0 && s.hourlyCounter >= s.maxPerHour {
|
||||
return false
|
||||
}
|
||||
|
||||
if s.maxPerDay > 0 && s.dailyCounter >= s.maxPerDay {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *AlphaVantageSource) incrementCounters() {
|
||||
s.hourlyCounter++
|
||||
s.dailyCounter++
|
||||
}
|
||||
|
||||
// FinnhubSource fetches news from Finnhub API
|
||||
type FinnhubSource struct {
|
||||
name string
|
||||
url string
|
||||
logger *slog.Logger
|
||||
httpClient *http.Client
|
||||
|
||||
// Rate limiting
|
||||
maxPerHour int
|
||||
maxPerDay int
|
||||
hourlyCounter int
|
||||
dailyCounter int
|
||||
lastHourReset time.Time
|
||||
lastDayReset time.Time
|
||||
}
|
||||
|
||||
type finnhubArticle struct {
|
||||
Category string `json:"category"`
|
||||
Datetime int64 `json:"datetime"`
|
||||
Headline string `json:"headline"`
|
||||
ID int64 `json:"id"`
|
||||
Image string `json:"image"`
|
||||
Related string `json:"related"`
|
||||
Source string `json:"source"`
|
||||
Summary string `json:"summary"`
|
||||
URL string `json:"url"`
|
||||
}
|
||||
|
||||
func NewFinnhubSource(name, url string, logger *slog.Logger) *FinnhubSource {
|
||||
now := time.Now()
|
||||
return &FinnhubSource{
|
||||
name: name,
|
||||
url: url,
|
||||
logger: logger,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
lastHourReset: now,
|
||||
lastDayReset: now,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *FinnhubSource) SetRateLimit(maxPerHour, maxPerDay int) {
|
||||
s.maxPerHour = maxPerHour
|
||||
s.maxPerDay = maxPerDay
|
||||
}
|
||||
|
||||
func (s *FinnhubSource) Name() string {
|
||||
return s.name
|
||||
}
|
||||
|
||||
func (s *FinnhubSource) Fetch(ctx context.Context) ([]*model.NewsArticle, error) {
|
||||
// Check rate limits
|
||||
if !s.checkRateLimit() {
|
||||
s.logger.Warn("rate limit reached, skipping fetch",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("hourly", s.hourlyCounter),
|
||||
slog.Int("max_per_hour", s.maxPerHour),
|
||||
slog.Int("daily", s.dailyCounter),
|
||||
slog.Int("max_per_day", s.maxPerDay))
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
s.logger.Debug("fetching Finnhub news", slog.String("source", s.name))
|
||||
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", s.url, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("http error: %s", resp.Status)
|
||||
}
|
||||
|
||||
var data []finnhubArticle
|
||||
if err := json.NewDecoder(resp.Body).Decode(&data); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse JSON: %w", err)
|
||||
}
|
||||
|
||||
var articles []*model.NewsArticle
|
||||
for _, item := range data {
|
||||
publishedAt := time.Unix(item.Datetime, 0)
|
||||
|
||||
// Extract symbols from Related field (comma-separated tickers)
|
||||
symbols := item.Related
|
||||
if symbols == "" {
|
||||
symbols = extractSymbols(item.Headline + " " + item.Summary)
|
||||
}
|
||||
|
||||
article := &model.NewsArticle{
|
||||
Source: s.name,
|
||||
Title: item.Headline,
|
||||
URL: item.URL,
|
||||
Content: item.Summary,
|
||||
PublishedAt: publishedAt,
|
||||
FetchedAt: time.Now(),
|
||||
Symbols: symbols,
|
||||
}
|
||||
|
||||
articles = append(articles, article)
|
||||
}
|
||||
|
||||
s.logger.Debug("fetched articles",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("count", len(articles)))
|
||||
|
||||
s.incrementCounters()
|
||||
return articles, nil
|
||||
}
|
||||
|
||||
func (s *FinnhubSource) checkRateLimit() bool {
|
||||
now := time.Now()
|
||||
|
||||
if now.Sub(s.lastHourReset) >= time.Hour {
|
||||
s.hourlyCounter = 0
|
||||
s.lastHourReset = now
|
||||
}
|
||||
|
||||
if now.Sub(s.lastDayReset) >= 24*time.Hour {
|
||||
s.dailyCounter = 0
|
||||
s.lastDayReset = now
|
||||
}
|
||||
|
||||
if s.maxPerHour > 0 && s.hourlyCounter >= s.maxPerHour {
|
||||
return false
|
||||
}
|
||||
|
||||
if s.maxPerDay > 0 && s.dailyCounter >= s.maxPerDay {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (s *FinnhubSource) incrementCounters() {
|
||||
s.hourlyCounter++
|
||||
s.dailyCounter++
|
||||
}
|
||||
@@ -0,0 +1,250 @@
|
||||
package news
|
||||
|
||||
import (
|
||||
"bytes"
|
||||
"context"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"io"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
// LLMScorer uses a local LLM (via Ollama) for contextual sentiment analysis
|
||||
type LLMScorer struct {
|
||||
cfg *config.LLMScorerConfig
|
||||
httpClient *http.Client
|
||||
fallbackAnalyzer *Analyzer
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
// LLMRequest represents the Ollama API request
|
||||
type LLMRequest struct {
|
||||
Model string `json:"model"`
|
||||
Prompt string `json:"prompt"`
|
||||
Temperature float64 `json:"temperature"`
|
||||
Stream bool `json:"stream"`
|
||||
}
|
||||
|
||||
// LLMResponse represents the Ollama API response
|
||||
type LLMResponse struct {
|
||||
Model string `json:"model"`
|
||||
CreatedAt string `json:"created_at"`
|
||||
Response string `json:"response"`
|
||||
Done bool `json:"done"`
|
||||
}
|
||||
|
||||
// SentimentResponse represents parsed sentiment from LLM
|
||||
type SentimentResponse struct {
|
||||
Sentiment string `json:"sentiment"` // "positive", "negative", "neutral"
|
||||
Score float64 `json:"score"` // -1.0 to 1.0
|
||||
Confidence float64 `json:"confidence"` // 0.0 to 1.0
|
||||
Reasoning string `json:"reasoning"`
|
||||
}
|
||||
|
||||
func NewLLMScorer(
|
||||
cfg *config.LLMScorerConfig,
|
||||
fallbackAnalyzer *Analyzer,
|
||||
logger *slog.Logger,
|
||||
) *LLMScorer {
|
||||
return &LLMScorer{
|
||||
cfg: cfg,
|
||||
httpClient: &http.Client{Timeout: cfg.Timeout.Duration},
|
||||
fallbackAnalyzer: fallbackAnalyzer,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (l *LLMScorer) Enabled() bool {
|
||||
return l.cfg.Enabled
|
||||
}
|
||||
|
||||
// Analyze uses LLM to analyze sentiment with optional ensemble mode
|
||||
func (l *LLMScorer) Analyze(ctx context.Context, article *model.NewsArticle) error {
|
||||
if !l.cfg.Enabled {
|
||||
return nil
|
||||
}
|
||||
|
||||
// Build structured prompt
|
||||
prompt := l.buildPrompt(article)
|
||||
|
||||
// Call Ollama API
|
||||
sentimentResp, err := l.callOllamaAPI(ctx, prompt)
|
||||
if err != nil {
|
||||
l.logger.Warn("LLM scoring failed, using keyword fallback",
|
||||
slog.String("article_id", fmt.Sprintf("%d", article.ID)),
|
||||
slog.Any("error", err))
|
||||
|
||||
// Fallback to keyword analyzer
|
||||
l.fallbackAnalyzer.Analyze(article)
|
||||
article.SentimentMethod = "keyword_fallback"
|
||||
return nil
|
||||
}
|
||||
|
||||
// Validate sentiment score
|
||||
if sentimentResp.Score < -1.0 || sentimentResp.Score > 1.0 {
|
||||
l.logger.Warn("LLM returned invalid score, using keyword fallback",
|
||||
slog.Float64("score", sentimentResp.Score))
|
||||
l.fallbackAnalyzer.Analyze(article)
|
||||
article.SentimentMethod = "keyword_fallback"
|
||||
return nil
|
||||
}
|
||||
|
||||
// Ensemble mode: weighted average of LLM + keyword scores
|
||||
if l.cfg.EnsembleWeight < 1.0 && l.fallbackAnalyzer != nil {
|
||||
// Get keyword score
|
||||
keywordArticle := &model.NewsArticle{
|
||||
Title: article.Title,
|
||||
Content: article.Content,
|
||||
}
|
||||
l.fallbackAnalyzer.Analyze(keywordArticle)
|
||||
|
||||
if keywordArticle.SentimentScore != nil {
|
||||
keywordScore := *keywordArticle.SentimentScore
|
||||
finalScore := (sentimentResp.Score * l.cfg.EnsembleWeight) +
|
||||
(keywordScore * (1.0 - l.cfg.EnsembleWeight))
|
||||
|
||||
article.SentimentScore = &finalScore
|
||||
article.SentimentMethod = "ensemble"
|
||||
|
||||
l.logger.Debug("ensemble scoring",
|
||||
slog.Float64("llm_score", sentimentResp.Score),
|
||||
slog.Float64("keyword_score", keywordScore),
|
||||
slog.Float64("final_score", finalScore),
|
||||
slog.Float64("weight", l.cfg.EnsembleWeight))
|
||||
} else {
|
||||
article.SentimentScore = &sentimentResp.Score
|
||||
article.SentimentMethod = "llm"
|
||||
}
|
||||
} else {
|
||||
// LLM only
|
||||
article.SentimentScore = &sentimentResp.Score
|
||||
article.SentimentMethod = "llm"
|
||||
}
|
||||
|
||||
// Set sentiment label based on final score
|
||||
if article.SentimentScore != nil {
|
||||
score := *article.SentimentScore
|
||||
if score > 0.3 {
|
||||
article.SentimentLabel = "positive"
|
||||
} else if score < -0.3 {
|
||||
article.SentimentLabel = "negative"
|
||||
} else {
|
||||
article.SentimentLabel = "neutral"
|
||||
}
|
||||
}
|
||||
|
||||
// Store LLM metadata
|
||||
article.LLMModel = &l.cfg.ModelName
|
||||
article.LLMConfidence = &sentimentResp.Confidence
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// buildPrompt creates a structured prompt for financial sentiment analysis
|
||||
func (l *LLMScorer) buildPrompt(article *model.NewsArticle) string {
|
||||
return fmt.Sprintf(`You are a financial sentiment analyzer. Respond ONLY with valid JSON (no markdown, no explanation).
|
||||
|
||||
Analyze this news article and determine if it's positive, negative, or neutral for stock trading:
|
||||
|
||||
---
|
||||
Title: %s
|
||||
Content: %s
|
||||
---
|
||||
|
||||
Respond with JSON in this exact format:
|
||||
{
|
||||
"sentiment": "positive" | "negative" | "neutral",
|
||||
"score": <float between -1.0 and 1.0>,
|
||||
"confidence": <float between 0.0 and 1.0>,
|
||||
"reasoning": "<brief one-sentence explanation>"
|
||||
}
|
||||
|
||||
JSON Response:`, article.Title, article.Content)
|
||||
}
|
||||
|
||||
// callOllamaAPI sends request to Ollama API and parses response
|
||||
func (l *LLMScorer) callOllamaAPI(ctx context.Context, prompt string) (*SentimentResponse, error) {
|
||||
// Build request
|
||||
reqBody := LLMRequest{
|
||||
Model: l.cfg.ModelName,
|
||||
Prompt: prompt,
|
||||
Temperature: l.cfg.Temperature,
|
||||
Stream: false,
|
||||
}
|
||||
|
||||
jsonData, err := json.Marshal(reqBody)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to marshal request: %w", err)
|
||||
}
|
||||
|
||||
// Create HTTP request with context
|
||||
req, err := http.NewRequestWithContext(ctx, "POST", l.cfg.Endpoint+"/api/generate", bytes.NewBuffer(jsonData))
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
req.Header.Set("Content-Type", "application/json")
|
||||
|
||||
// Send request
|
||||
resp, err := l.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to call Ollama API: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
body, _ := io.ReadAll(resp.Body)
|
||||
return nil, fmt.Errorf("Ollama API returned status %d: %s", resp.StatusCode, string(body))
|
||||
}
|
||||
|
||||
// Read response
|
||||
body, err := io.ReadAll(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read response: %w", err)
|
||||
}
|
||||
|
||||
var llmResp LLMResponse
|
||||
if err := json.Unmarshal(body, &llmResp); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse Ollama response: %w", err)
|
||||
}
|
||||
|
||||
// Parse sentiment from LLM response
|
||||
return l.parseSentiment(llmResp.Response)
|
||||
}
|
||||
|
||||
// parseSentiment extracts structured sentiment from LLM text response
|
||||
func (l *LLMScorer) parseSentiment(response string) (*SentimentResponse, error) {
|
||||
// Try to find JSON in response (LLM might add text before/after)
|
||||
startIdx := -1
|
||||
endIdx := -1
|
||||
|
||||
for i := 0; i < len(response); i++ {
|
||||
if response[i] == '{' && startIdx == -1 {
|
||||
startIdx = i
|
||||
}
|
||||
if response[i] == '}' {
|
||||
endIdx = i + 1
|
||||
}
|
||||
}
|
||||
|
||||
if startIdx == -1 || endIdx == -1 {
|
||||
return nil, fmt.Errorf("no JSON found in LLM response")
|
||||
}
|
||||
|
||||
jsonStr := response[startIdx:endIdx]
|
||||
|
||||
var sentiment SentimentResponse
|
||||
if err := json.Unmarshal([]byte(jsonStr), &sentiment); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse sentiment JSON: %w", err)
|
||||
}
|
||||
|
||||
// Validate sentiment field
|
||||
if sentiment.Sentiment != "positive" && sentiment.Sentiment != "negative" && sentiment.Sentiment != "neutral" {
|
||||
return nil, fmt.Errorf("invalid sentiment value: %s", sentiment.Sentiment)
|
||||
}
|
||||
|
||||
return &sentiment, nil
|
||||
}
|
||||
@@ -0,0 +1,222 @@
|
||||
package news
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/mmcdole/gofeed"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type Source interface {
|
||||
Name() string
|
||||
Fetch(ctx context.Context) ([]*model.NewsArticle, error)
|
||||
}
|
||||
|
||||
// RSSSource fetches news from RSS feeds
|
||||
type RSSSource struct {
|
||||
name string
|
||||
url string
|
||||
parser *gofeed.Parser
|
||||
logger *slog.Logger
|
||||
httpClient *http.Client
|
||||
headers map[string]string
|
||||
username string
|
||||
password string
|
||||
|
||||
// Rate limiting
|
||||
maxPerHour int
|
||||
maxPerDay int
|
||||
hourlyCounter int
|
||||
dailyCounter int
|
||||
lastHourReset time.Time
|
||||
lastDayReset time.Time
|
||||
}
|
||||
|
||||
func NewRSSSource(name, url string, logger *slog.Logger) *RSSSource {
|
||||
now := time.Now()
|
||||
return &RSSSource{
|
||||
name: name,
|
||||
url: url,
|
||||
parser: gofeed.NewParser(),
|
||||
logger: logger,
|
||||
httpClient: &http.Client{Timeout: 30 * time.Second},
|
||||
headers: make(map[string]string),
|
||||
lastHourReset: now,
|
||||
lastDayReset: now,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *RSSSource) SetRateLimit(maxPerHour, maxPerDay int) {
|
||||
s.maxPerHour = maxPerHour
|
||||
s.maxPerDay = maxPerDay
|
||||
}
|
||||
|
||||
func (s *RSSSource) SetBasicAuth(username, password string) {
|
||||
s.username = username
|
||||
s.password = password
|
||||
}
|
||||
|
||||
func (s *RSSSource) AddHeader(key, value string) {
|
||||
s.headers[key] = value
|
||||
}
|
||||
|
||||
func (s *RSSSource) Name() string {
|
||||
return s.name
|
||||
}
|
||||
|
||||
func (s *RSSSource) Fetch(ctx context.Context) ([]*model.NewsArticle, error) {
|
||||
// Check rate limits
|
||||
if !s.checkRateLimit() {
|
||||
s.logger.Warn("rate limit reached, skipping fetch",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("hourly", s.hourlyCounter),
|
||||
slog.Int("max_per_hour", s.maxPerHour),
|
||||
slog.Int("daily", s.dailyCounter),
|
||||
slog.Int("max_per_day", s.maxPerDay))
|
||||
return nil, nil // Return empty, not an error
|
||||
}
|
||||
|
||||
s.logger.Debug("fetching RSS feed", slog.String("source", s.name), slog.String("url", s.url))
|
||||
|
||||
// Create HTTP request with custom headers and auth
|
||||
req, err := http.NewRequestWithContext(ctx, "GET", s.url, nil)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create request: %w", err)
|
||||
}
|
||||
|
||||
// Add custom headers
|
||||
for key, value := range s.headers {
|
||||
req.Header.Set(key, value)
|
||||
}
|
||||
|
||||
// Add basic auth if configured
|
||||
if s.username != "" {
|
||||
req.SetBasicAuth(s.username, s.password)
|
||||
}
|
||||
|
||||
// Fetch the feed
|
||||
resp, err := s.httpClient.Do(req)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to fetch feed: %w", err)
|
||||
}
|
||||
defer resp.Body.Close()
|
||||
|
||||
if resp.StatusCode != http.StatusOK {
|
||||
return nil, fmt.Errorf("http error: %s", resp.Status)
|
||||
}
|
||||
|
||||
// Parse the feed
|
||||
feed, err := s.parser.Parse(resp.Body)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse RSS feed: %w", err)
|
||||
}
|
||||
|
||||
var articles []*model.NewsArticle
|
||||
for _, item := range feed.Items {
|
||||
publishedAt := time.Now()
|
||||
if item.PublishedParsed != nil {
|
||||
publishedAt = *item.PublishedParsed
|
||||
} else if item.UpdatedParsed != nil {
|
||||
publishedAt = *item.UpdatedParsed
|
||||
}
|
||||
|
||||
content := item.Description
|
||||
if item.Content != "" {
|
||||
content = item.Content
|
||||
}
|
||||
|
||||
article := &model.NewsArticle{
|
||||
Source: s.name,
|
||||
Title: item.Title,
|
||||
URL: item.Link,
|
||||
Content: content,
|
||||
PublishedAt: publishedAt,
|
||||
FetchedAt: time.Now(),
|
||||
Symbols: extractSymbols(item.Title + " " + content),
|
||||
}
|
||||
|
||||
articles = append(articles, article)
|
||||
}
|
||||
|
||||
s.logger.Debug("fetched articles",
|
||||
slog.String("source", s.name),
|
||||
slog.Int("count", len(articles)),
|
||||
)
|
||||
|
||||
// Increment rate limit counters
|
||||
s.incrementCounters()
|
||||
|
||||
return articles, nil
|
||||
}
|
||||
|
||||
// checkRateLimit checks if we can make a request based on rate limits
|
||||
func (s *RSSSource) checkRateLimit() bool {
|
||||
now := time.Now()
|
||||
|
||||
// Reset hourly counter if an hour has passed
|
||||
if now.Sub(s.lastHourReset) >= time.Hour {
|
||||
s.hourlyCounter = 0
|
||||
s.lastHourReset = now
|
||||
}
|
||||
|
||||
// Reset daily counter if a day has passed
|
||||
if now.Sub(s.lastDayReset) >= 24*time.Hour {
|
||||
s.dailyCounter = 0
|
||||
s.lastDayReset = now
|
||||
}
|
||||
|
||||
// Check hourly limit (0 means unlimited)
|
||||
if s.maxPerHour > 0 && s.hourlyCounter >= s.maxPerHour {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check daily limit (0 means unlimited)
|
||||
if s.maxPerDay > 0 && s.dailyCounter >= s.maxPerDay {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
// incrementCounters increments the rate limit counters after a successful fetch
|
||||
func (s *RSSSource) incrementCounters() {
|
||||
s.hourlyCounter++
|
||||
s.dailyCounter++
|
||||
}
|
||||
|
||||
// extractSymbols extracts potential stock symbols from text
|
||||
// Simple implementation - looks for uppercase words 1-5 chars long
|
||||
func extractSymbols(text string) string {
|
||||
words := strings.Fields(text)
|
||||
var symbols []string
|
||||
seen := make(map[string]bool)
|
||||
|
||||
for _, word := range words {
|
||||
// Clean word
|
||||
word = strings.Trim(word, ".,!?;:\"'()")
|
||||
|
||||
// Check if it looks like a stock symbol
|
||||
if len(word) >= 1 && len(word) <= 5 && isAllUppercase(word) {
|
||||
if !seen[word] {
|
||||
symbols = append(symbols, word)
|
||||
seen[word] = true
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return strings.Join(symbols, ",")
|
||||
}
|
||||
|
||||
func isAllUppercase(s string) bool {
|
||||
for _, r := range s {
|
||||
if r < 'A' || r > 'Z' {
|
||||
return false
|
||||
}
|
||||
}
|
||||
return len(s) > 0
|
||||
}
|
||||
@@ -0,0 +1,114 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type AggressiveStrategy struct {
|
||||
stopLossEnabled bool
|
||||
stopLossPercent float64
|
||||
}
|
||||
|
||||
func NewAggressiveStrategy(stopLossEnabled bool, stopLossPercent float64) *AggressiveStrategy {
|
||||
return &AggressiveStrategy{
|
||||
stopLossEnabled: stopLossEnabled,
|
||||
stopLossPercent: stopLossPercent,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AggressiveStrategy) Name() string {
|
||||
return "aggressive"
|
||||
}
|
||||
|
||||
func (s *AggressiveStrategy) GetRiskParams() RiskParameters {
|
||||
return RiskParameters{
|
||||
MaxParallelTrades: 10,
|
||||
MaxTradesPerHour: 12,
|
||||
PositionSizePercent: 7.5, // 5-10% of capital
|
||||
StopLossPercent: 5.0,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *AggressiveStrategy) Analyze(ctx context.Context, market *MarketData, news []*model.NewsArticle) (*TradeSignal, error) {
|
||||
// Aggressive strategy: Trade on any positive sentiment
|
||||
// Higher risk, higher frequency
|
||||
|
||||
positiveCount := 0
|
||||
negativeCount := 0
|
||||
totalSentiment := 0.0
|
||||
relevantArticles := 0
|
||||
|
||||
for _, article := range news {
|
||||
if article.SentimentScore == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if article mentions this symbol
|
||||
if !strings.Contains(strings.ToUpper(article.Symbols), market.Symbol) {
|
||||
continue
|
||||
}
|
||||
|
||||
relevantArticles++
|
||||
sentiment := *article.SentimentScore
|
||||
totalSentiment += sentiment
|
||||
|
||||
if article.SentimentLabel == "positive" {
|
||||
positiveCount++
|
||||
} else if article.SentimentLabel == "negative" {
|
||||
negativeCount++
|
||||
}
|
||||
}
|
||||
|
||||
// Aggressive: Need at least 1 relevant article
|
||||
if relevantArticles < 1 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
avgSentiment := totalSentiment / float64(relevantArticles)
|
||||
|
||||
// Aggressive: Any net positive sentiment
|
||||
if avgSentiment <= 0 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Also consider sell signals on strong negative sentiment
|
||||
if avgSentiment < -0.5 && negativeCount > positiveCount {
|
||||
return &TradeSignal{
|
||||
Symbol: market.Symbol,
|
||||
Action: model.ActionSell,
|
||||
Quantity: 0, // Will be filled from position
|
||||
Confidence: (-avgSentiment) * 0.95,
|
||||
Reasoning: fmt.Sprintf("Aggressive SELL: %d negative articles, avg sentiment %.2f", negativeCount, avgSentiment),
|
||||
}, nil
|
||||
}
|
||||
|
||||
return &TradeSignal{
|
||||
Symbol: market.Symbol,
|
||||
Action: model.ActionBuy,
|
||||
Quantity: 0, // Will be calculated by trader
|
||||
Confidence: avgSentiment * 0.95,
|
||||
Reasoning: fmt.Sprintf("Aggressive BUY: %d positive articles, avg sentiment %.2f", positiveCount, avgSentiment),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CalculatePositionSize - Aggressive strategy uses larger sizing with confidence scaling
|
||||
func (s *AggressiveStrategy) CalculatePositionSize(price float64, confidence float64, availableCapital float64) int {
|
||||
basePositionPercent := s.GetRiskParams().PositionSizePercent
|
||||
|
||||
// Aggressive: More aggressive confidence scaling (0.7 - 1.2x)
|
||||
confidenceMultiplier := 0.7 + (confidence * 0.5)
|
||||
adjustedPercent := basePositionPercent * confidenceMultiplier
|
||||
|
||||
positionValue := availableCapital * (adjustedPercent / 100.0)
|
||||
quantity := int(positionValue / price)
|
||||
|
||||
if quantity < 1 {
|
||||
return 1
|
||||
}
|
||||
|
||||
return quantity
|
||||
}
|
||||
@@ -0,0 +1,101 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type DefensiveStrategy struct {
|
||||
stopLossEnabled bool
|
||||
stopLossPercent float64
|
||||
}
|
||||
|
||||
func NewDefensiveStrategy(stopLossEnabled bool, stopLossPercent float64) *DefensiveStrategy {
|
||||
return &DefensiveStrategy{
|
||||
stopLossEnabled: stopLossEnabled,
|
||||
stopLossPercent: stopLossPercent,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *DefensiveStrategy) Name() string {
|
||||
return "defensive"
|
||||
}
|
||||
|
||||
func (s *DefensiveStrategy) GetRiskParams() RiskParameters {
|
||||
return RiskParameters{
|
||||
MaxParallelTrades: 2,
|
||||
MaxTradesPerHour: 3,
|
||||
PositionSizePercent: 1.5, // 1-2% of capital
|
||||
StopLossPercent: 2.0,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *DefensiveStrategy) Analyze(ctx context.Context, market *MarketData, news []*model.NewsArticle) (*TradeSignal, error) {
|
||||
// Defensive strategy: Only trade on strong positive sentiment
|
||||
// with multiple confirming news articles
|
||||
|
||||
positiveCount := 0
|
||||
negativeCount := 0
|
||||
totalSentiment := 0.0
|
||||
|
||||
for _, article := range news {
|
||||
if article.SentimentScore == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if article mentions this symbol
|
||||
if !strings.Contains(strings.ToUpper(article.Symbols), market.Symbol) {
|
||||
continue
|
||||
}
|
||||
|
||||
sentiment := *article.SentimentScore
|
||||
totalSentiment += sentiment
|
||||
|
||||
if article.SentimentLabel == "positive" {
|
||||
positiveCount++
|
||||
} else if article.SentimentLabel == "negative" {
|
||||
negativeCount++
|
||||
}
|
||||
}
|
||||
|
||||
// Defensive: Need at least 3 positive articles and no negative ones
|
||||
if positiveCount < 3 || negativeCount > 0 {
|
||||
return nil, nil // No trade signal
|
||||
}
|
||||
|
||||
avgSentiment := totalSentiment / float64(len(news))
|
||||
|
||||
// Strong positive sentiment required (>0.5)
|
||||
if avgSentiment < 0.5 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return &TradeSignal{
|
||||
Symbol: market.Symbol,
|
||||
Action: model.ActionBuy,
|
||||
Quantity: 0, // Will be calculated by trader
|
||||
Confidence: avgSentiment * 0.8, // Conservative confidence
|
||||
Reasoning: fmt.Sprintf("Defensive: %d positive news articles, avg sentiment %.2f", positiveCount, avgSentiment),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CalculatePositionSize - Defensive strategy uses conservative sizing with confidence scaling
|
||||
func (s *DefensiveStrategy) CalculatePositionSize(price float64, confidence float64, availableCapital float64) int {
|
||||
basePositionPercent := s.GetRiskParams().PositionSizePercent
|
||||
|
||||
// Defensive: More conservative confidence scaling (0.3 - 0.8x)
|
||||
confidenceMultiplier := 0.3 + (confidence * 0.5)
|
||||
adjustedPercent := basePositionPercent * confidenceMultiplier
|
||||
|
||||
positionValue := availableCapital * (adjustedPercent / 100.0)
|
||||
quantity := int(positionValue / price)
|
||||
|
||||
if quantity < 1 {
|
||||
return 1
|
||||
}
|
||||
|
||||
return quantity
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"strings"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type NormalStrategy struct {
|
||||
stopLossEnabled bool
|
||||
stopLossPercent float64
|
||||
}
|
||||
|
||||
func NewNormalStrategy(stopLossEnabled bool, stopLossPercent float64) *NormalStrategy {
|
||||
return &NormalStrategy{
|
||||
stopLossEnabled: stopLossEnabled,
|
||||
stopLossPercent: stopLossPercent,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *NormalStrategy) Name() string {
|
||||
return "normal"
|
||||
}
|
||||
|
||||
func (s *NormalStrategy) GetRiskParams() RiskParameters {
|
||||
return RiskParameters{
|
||||
MaxParallelTrades: 5,
|
||||
MaxTradesPerHour: 6,
|
||||
PositionSizePercent: 4.0, // 3-5% of capital
|
||||
StopLossPercent: 3.0,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *NormalStrategy) Analyze(ctx context.Context, market *MarketData, news []*model.NewsArticle) (*TradeSignal, error) {
|
||||
// Normal strategy: Trade on moderate positive sentiment for BUY
|
||||
// Sell on negative sentiment or profit target
|
||||
|
||||
positiveCount := 0
|
||||
negativeCount := 0
|
||||
totalSentiment := 0.0
|
||||
relevantArticles := 0
|
||||
|
||||
for _, article := range news {
|
||||
if article.SentimentScore == nil {
|
||||
continue
|
||||
}
|
||||
|
||||
// Check if article mentions this symbol
|
||||
if !strings.Contains(strings.ToUpper(article.Symbols), market.Symbol) {
|
||||
continue
|
||||
}
|
||||
|
||||
relevantArticles++
|
||||
sentiment := *article.SentimentScore
|
||||
totalSentiment += sentiment
|
||||
|
||||
if article.SentimentLabel == "positive" {
|
||||
positiveCount++
|
||||
} else if article.SentimentLabel == "negative" {
|
||||
negativeCount++
|
||||
}
|
||||
}
|
||||
|
||||
// Need at least 2 relevant articles
|
||||
if relevantArticles < 2 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
avgSentiment := totalSentiment / float64(relevantArticles)
|
||||
|
||||
// SELL signal: Strong negative sentiment
|
||||
if negativeCount > positiveCount && avgSentiment < -0.3 {
|
||||
quantity := 0 // Will be filled from position
|
||||
|
||||
return &TradeSignal{
|
||||
Symbol: market.Symbol,
|
||||
Action: model.ActionSell,
|
||||
Quantity: quantity,
|
||||
Confidence: -avgSentiment * 0.9, // Convert negative to positive confidence
|
||||
Reasoning: fmt.Sprintf("Normal SELL: %d negative vs %d positive articles, avg sentiment %.2f", negativeCount, positiveCount, avgSentiment),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// BUY signal: Positive sentiment outweighs negative
|
||||
if positiveCount <= negativeCount {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
// Moderate positive sentiment required (>0.3)
|
||||
if avgSentiment < 0.3 {
|
||||
return nil, nil
|
||||
}
|
||||
|
||||
return &TradeSignal{
|
||||
Symbol: market.Symbol,
|
||||
Action: model.ActionBuy,
|
||||
Quantity: 0, // Will be calculated by trader with current balance
|
||||
Confidence: avgSentiment * 0.9,
|
||||
Reasoning: fmt.Sprintf("Normal BUY: %d positive vs %d negative articles, avg sentiment %.2f", positiveCount, negativeCount, avgSentiment),
|
||||
}, nil
|
||||
}
|
||||
|
||||
// CalculatePositionSize determines how many shares to buy based on confidence and available capital
|
||||
func (s *NormalStrategy) CalculatePositionSize(price float64, confidence float64, availableCapital float64) int {
|
||||
// Base position size from strategy risk params
|
||||
basePositionPercent := s.GetRiskParams().PositionSizePercent
|
||||
|
||||
// Scale position size by confidence (0.5 - 1.0 confidence → 0.5x - 1.0x of base)
|
||||
// High confidence = larger position, low confidence = smaller position
|
||||
confidenceMultiplier := 0.5 + (confidence * 0.5)
|
||||
adjustedPercent := basePositionPercent * confidenceMultiplier
|
||||
|
||||
// Calculate position value and quantity
|
||||
positionValue := availableCapital * (adjustedPercent / 100.0)
|
||||
quantity := int(positionValue / price)
|
||||
|
||||
if quantity < 1 {
|
||||
return 1 // Minimum 1 share
|
||||
}
|
||||
|
||||
return quantity
|
||||
}
|
||||
|
||||
// ValidateTradeValue checks if trade value is within absolute maximum
|
||||
func ValidateTradeValue(quantity int, price float64, maxTradeValue float64) (int, error) {
|
||||
if maxTradeValue <= 0 {
|
||||
return quantity, nil // No limit
|
||||
}
|
||||
|
||||
tradeValue := float64(quantity) * price
|
||||
|
||||
if tradeValue <= maxTradeValue {
|
||||
return quantity, nil // Within limit
|
||||
}
|
||||
|
||||
// Calculate max quantity that fits within limit
|
||||
maxQuantity := int(maxTradeValue / price)
|
||||
|
||||
if maxQuantity < 1 {
|
||||
return 0, fmt.Errorf("trade value would be $%.2f but max is $%.2f (price $%.2f too high for 1 share)",
|
||||
tradeValue, maxTradeValue, price)
|
||||
}
|
||||
|
||||
return maxQuantity, nil
|
||||
}
|
||||
@@ -0,0 +1,51 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type Strategy interface {
|
||||
Name() string
|
||||
Analyze(ctx context.Context, market *MarketData, news []*model.NewsArticle) (*TradeSignal, error)
|
||||
CalculatePositionSize(price float64, confidence float64, availableCapital float64) int
|
||||
GetRiskParams() RiskParameters
|
||||
}
|
||||
|
||||
type MarketData struct {
|
||||
Symbol string
|
||||
LastPrice float64
|
||||
BidPrice float64
|
||||
AskPrice float64
|
||||
Volume int64
|
||||
Change float64
|
||||
ChangePercent float64
|
||||
}
|
||||
|
||||
type TradeSignal struct {
|
||||
Symbol string
|
||||
Action model.ActionType
|
||||
Quantity int
|
||||
Confidence float64 // 0.0 - 1.0
|
||||
Reasoning string
|
||||
}
|
||||
|
||||
type RiskParameters struct {
|
||||
MaxParallelTrades int
|
||||
MaxTradesPerHour int
|
||||
PositionSizePercent float64 // Percentage of capital per trade
|
||||
StopLossPercent float64
|
||||
}
|
||||
|
||||
// Factory function to create strategy based on name
|
||||
func NewStrategy(strategyName string, stopLossEnabled bool, stopLossPercent float64) Strategy {
|
||||
switch strategyName {
|
||||
case "defensive":
|
||||
return NewDefensiveStrategy(stopLossEnabled, stopLossPercent)
|
||||
case "aggressive":
|
||||
return NewAggressiveStrategy(stopLossEnabled, stopLossPercent)
|
||||
default: // "normal"
|
||||
return NewNormalStrategy(stopLossEnabled, stopLossPercent)
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,123 @@
|
||||
package strategy
|
||||
|
||||
import (
|
||||
"context"
|
||||
"testing"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
func TestDefensiveStrategy(t *testing.T) {
|
||||
strategy := NewDefensiveStrategy(true, 2.0)
|
||||
|
||||
if strategy.Name() != "defensive" {
|
||||
t.Errorf("Expected name 'defensive', got '%s'", strategy.Name())
|
||||
}
|
||||
|
||||
params := strategy.GetRiskParams()
|
||||
if params.MaxParallelTrades != 2 {
|
||||
t.Errorf("Expected MaxParallelTrades=2, got %d", params.MaxParallelTrades)
|
||||
}
|
||||
if params.MaxTradesPerHour != 3 {
|
||||
t.Errorf("Expected MaxTradesPerHour=3, got %d", params.MaxTradesPerHour)
|
||||
}
|
||||
}
|
||||
|
||||
func TestNormalStrategy(t *testing.T) {
|
||||
strategy := NewNormalStrategy(true, 3.0)
|
||||
|
||||
if strategy.Name() != "normal" {
|
||||
t.Errorf("Expected name 'normal', got '%s'", strategy.Name())
|
||||
}
|
||||
|
||||
params := strategy.GetRiskParams()
|
||||
if params.MaxParallelTrades != 5 {
|
||||
t.Errorf("Expected MaxParallelTrades=5, got %d", params.MaxParallelTrades)
|
||||
}
|
||||
if params.MaxTradesPerHour != 6 {
|
||||
t.Errorf("Expected MaxTradesPerHour=6, got %d", params.MaxTradesPerHour)
|
||||
}
|
||||
}
|
||||
|
||||
func TestAggressiveStrategy(t *testing.T) {
|
||||
strategy := NewAggressiveStrategy(true, 5.0)
|
||||
|
||||
if strategy.Name() != "aggressive" {
|
||||
t.Errorf("Expected name 'aggressive', got '%s'", strategy.Name())
|
||||
}
|
||||
|
||||
params := strategy.GetRiskParams()
|
||||
if params.MaxParallelTrades != 10 {
|
||||
t.Errorf("Expected MaxParallelTrades=10, got %d", params.MaxParallelTrades)
|
||||
}
|
||||
if params.MaxTradesPerHour != 12 {
|
||||
t.Errorf("Expected MaxTradesPerHour=12, got %d", params.MaxTradesPerHour)
|
||||
}
|
||||
}
|
||||
|
||||
func TestStrategyFactory(t *testing.T) {
|
||||
defensive := NewStrategy("defensive", true, 2.0)
|
||||
if defensive.Name() != "defensive" {
|
||||
t.Error("Factory should create defensive strategy")
|
||||
}
|
||||
|
||||
normal := NewStrategy("normal", true, 3.0)
|
||||
if normal.Name() != "normal" {
|
||||
t.Error("Factory should create normal strategy")
|
||||
}
|
||||
|
||||
aggressive := NewStrategy("aggressive", true, 5.0)
|
||||
if aggressive.Name() != "aggressive" {
|
||||
t.Error("Factory should create aggressive strategy")
|
||||
}
|
||||
|
||||
// Default to normal
|
||||
defaultStrat := NewStrategy("unknown", true, 3.0)
|
||||
if defaultStrat.Name() != "normal" {
|
||||
t.Error("Factory should default to normal strategy")
|
||||
}
|
||||
}
|
||||
|
||||
func TestAnalyzeWithPositiveSentiment(t *testing.T) {
|
||||
strategy := NewNormalStrategy(true, 3.0)
|
||||
|
||||
market := &MarketData{
|
||||
Symbol: "AAPL",
|
||||
LastPrice: 150.0,
|
||||
}
|
||||
|
||||
sentiment := 0.8
|
||||
news := []*model.NewsArticle{
|
||||
{
|
||||
Title: "AAPL shows strong growth",
|
||||
Symbols: "AAPL",
|
||||
SentimentScore: &sentiment,
|
||||
SentimentLabel: "positive",
|
||||
},
|
||||
{
|
||||
Title: "Apple AAPL beats expectations",
|
||||
Symbols: "AAPL",
|
||||
SentimentScore: &sentiment,
|
||||
SentimentLabel: "positive",
|
||||
},
|
||||
}
|
||||
|
||||
signal, err := strategy.Analyze(context.Background(), market, news)
|
||||
if err != nil {
|
||||
t.Fatalf("Analyze failed: %v", err)
|
||||
}
|
||||
|
||||
if signal == nil {
|
||||
t.Error("Expected trade signal, got nil")
|
||||
} else {
|
||||
if signal.Symbol != "AAPL" {
|
||||
t.Errorf("Expected symbol AAPL, got %s", signal.Symbol)
|
||||
}
|
||||
if signal.Action != model.ActionBuy {
|
||||
t.Errorf("Expected BUY action, got %s", signal.Action)
|
||||
}
|
||||
if signal.Confidence <= 0 || signal.Confidence > 1 {
|
||||
t.Errorf("Confidence should be between 0 and 1, got %.2f", signal.Confidence)
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,205 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math/rand"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type DryRunExecutor struct {
|
||||
tradeRepo *db.TradeRepository
|
||||
balanceRepo *db.BalanceRepository
|
||||
positionRepo *db.PositionRepository
|
||||
logger *slog.Logger
|
||||
balance float64
|
||||
positionPrices map[string]float64 // Track entry prices for P&L calculation
|
||||
}
|
||||
|
||||
func NewDryRunExecutor(
|
||||
tradeRepo *db.TradeRepository,
|
||||
balanceRepo *db.BalanceRepository,
|
||||
positionRepo *db.PositionRepository,
|
||||
startingBalance float64,
|
||||
logger *slog.Logger,
|
||||
) *DryRunExecutor {
|
||||
return &DryRunExecutor{
|
||||
tradeRepo: tradeRepo,
|
||||
balanceRepo: balanceRepo,
|
||||
positionRepo: positionRepo,
|
||||
logger: logger,
|
||||
balance: startingBalance,
|
||||
positionPrices: make(map[string]float64),
|
||||
}
|
||||
}
|
||||
|
||||
func (e *DryRunExecutor) ExecuteTrade(ctx context.Context, trade *model.Trade) error {
|
||||
e.logger.Info("executing dry run trade",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.String("action", string(trade.Action)),
|
||||
slog.Int("quantity", trade.Quantity),
|
||||
slog.Float64("current_balance", e.balance),
|
||||
)
|
||||
|
||||
// Simulate order execution with random price movement
|
||||
simulatedPrice := e.simulatePrice(trade)
|
||||
|
||||
// Calculate cost
|
||||
cost := simulatedPrice * float64(trade.Quantity)
|
||||
|
||||
// Check if we have enough balance for BUY
|
||||
if trade.Action == model.ActionBuy {
|
||||
if cost > e.balance {
|
||||
return fmt.Errorf("insufficient dry run balance: need %.2f, have %.2f", cost, e.balance)
|
||||
}
|
||||
e.balance -= cost
|
||||
e.positionPrices[trade.Symbol] = simulatedPrice
|
||||
|
||||
// Create position record
|
||||
position := &model.Position{
|
||||
Symbol: trade.Symbol,
|
||||
Quantity: trade.Quantity,
|
||||
EntryPrice: simulatedPrice,
|
||||
EntryTradeID: trade.ID,
|
||||
}
|
||||
if err := e.positionRepo.Create(ctx, position); err != nil {
|
||||
e.logger.Error("failed to create position", slog.Any("error", err))
|
||||
}
|
||||
|
||||
e.logger.Info("dry run BUY executed",
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.Float64("price", simulatedPrice),
|
||||
slog.Float64("cost", cost),
|
||||
slog.Float64("remaining_balance", e.balance),
|
||||
)
|
||||
} else {
|
||||
// SELL - calculate P&L
|
||||
entryPrice, exists := e.positionPrices[trade.Symbol]
|
||||
if !exists {
|
||||
entryPrice = simulatedPrice * 0.95 // Assume we bought 5% lower
|
||||
}
|
||||
|
||||
pnl := (simulatedPrice - entryPrice) * float64(trade.Quantity)
|
||||
e.balance += cost
|
||||
trade.DryRunPnL = &pnl
|
||||
|
||||
delete(e.positionPrices, trade.Symbol)
|
||||
|
||||
// Remove position record
|
||||
position, err := e.positionRepo.GetBySymbol(ctx, trade.Symbol)
|
||||
if err == nil && position != nil {
|
||||
if err := e.positionRepo.Delete(ctx, position.ID); err != nil {
|
||||
e.logger.Error("failed to delete position", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
|
||||
e.logger.Info("dry run SELL executed",
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.Float64("entry_price", entryPrice),
|
||||
slog.Float64("exit_price", simulatedPrice),
|
||||
slog.Float64("pnl", pnl),
|
||||
slog.Float64("new_balance", e.balance),
|
||||
)
|
||||
}
|
||||
|
||||
// Update trade status
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeSubmitted
|
||||
trade.SubmittedAt = &now
|
||||
trade.ExecutedPrice = &simulatedPrice
|
||||
trade.IsDryRun = true
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update trade status: %w", err)
|
||||
}
|
||||
|
||||
// Simulate fill after 1 second
|
||||
go e.simulateFill(trade)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *DryRunExecutor) simulatePrice(trade *model.Trade) float64 {
|
||||
// Use target price if available, otherwise simulate
|
||||
if trade.TargetPrice != nil {
|
||||
// Add some random slippage (-0.5% to +0.5%)
|
||||
slippage := (*trade.TargetPrice) * (rand.Float64() - 0.5) * 0.01
|
||||
return *trade.TargetPrice + slippage
|
||||
}
|
||||
|
||||
// Generate a random price based on symbol
|
||||
// In production, you'd fetch real market data
|
||||
basePrice := 150.0
|
||||
|
||||
// Add some randomness
|
||||
movement := (rand.Float64() - 0.5) * 10.0
|
||||
return basePrice + movement
|
||||
}
|
||||
|
||||
func (e *DryRunExecutor) simulateFill(trade *model.Trade) {
|
||||
time.Sleep(1 * time.Second)
|
||||
|
||||
ctx := context.Background()
|
||||
now := time.Now()
|
||||
|
||||
trade.Status = model.TradeFilled
|
||||
trade.FilledAt = &now
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
e.logger.Error("failed to update filled trade", slog.Any("error", err))
|
||||
return
|
||||
}
|
||||
|
||||
e.logger.Info("dry run trade filled",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Float64("executed_price", *trade.ExecutedPrice),
|
||||
)
|
||||
|
||||
// Mark as completed (no stop-loss in dry run for simplicity)
|
||||
time.Sleep(500 * time.Millisecond)
|
||||
|
||||
now = time.Now()
|
||||
trade.Status = model.TradeCompleted
|
||||
trade.CompletedAt = &now
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
e.logger.Error("failed to mark trade completed", slog.Any("error", err))
|
||||
return
|
||||
}
|
||||
|
||||
// Update balance in database
|
||||
balance := &model.Balance{
|
||||
Timestamp: time.Now(),
|
||||
TotalValue: e.balance,
|
||||
CashBalance: e.balance,
|
||||
BuyingPower: e.balance,
|
||||
UnrealizedPnL: nil,
|
||||
RealizedPnL: trade.DryRunPnL,
|
||||
}
|
||||
|
||||
if err := e.balanceRepo.Create(ctx, balance); err != nil {
|
||||
e.logger.Error("failed to store dry run balance", slog.Any("error", err))
|
||||
return
|
||||
}
|
||||
|
||||
e.logger.Info("dry run trade completed",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Float64("balance", e.balance),
|
||||
)
|
||||
}
|
||||
|
||||
func (e *DryRunExecutor) GetBalance() float64 {
|
||||
return e.balance
|
||||
}
|
||||
|
||||
func (e *DryRunExecutor) CancelTrade(ctx context.Context, trade *model.Trade) error {
|
||||
e.logger.Info("dry run trade cancelled",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
)
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/app/client"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type Executor struct {
|
||||
client *client.IBClient
|
||||
tradeRepo *db.TradeRepository
|
||||
stopLossMgr *StopLossManager
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func NewExecutor(
|
||||
client *client.IBClient,
|
||||
tradeRepo *db.TradeRepository,
|
||||
stopLossMgr *StopLossManager,
|
||||
logger *slog.Logger,
|
||||
) *Executor {
|
||||
return &Executor{
|
||||
client: client,
|
||||
tradeRepo: tradeRepo,
|
||||
stopLossMgr: stopLossMgr,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (e *Executor) ExecuteTrade(ctx context.Context, trade *model.Trade) error {
|
||||
e.logger.Info("executing trade",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.String("action", string(trade.Action)),
|
||||
slog.Int("quantity", trade.Quantity),
|
||||
)
|
||||
|
||||
// Create order
|
||||
order := &client.Order{
|
||||
Symbol: trade.Symbol,
|
||||
Action: trade.Action,
|
||||
Quantity: trade.Quantity,
|
||||
OrderType: "MKT", // Market order
|
||||
}
|
||||
|
||||
// Place order with IB
|
||||
orderID, err := e.client.PlaceOrder(ctx, order)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to place order: %w", err)
|
||||
}
|
||||
|
||||
// Update trade status
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeSubmitted
|
||||
trade.SubmittedAt = &now
|
||||
trade.IBOrderID = &orderID
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update trade status: %w", err)
|
||||
}
|
||||
|
||||
e.logger.Info("trade submitted to IB",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Int64("ib_order_id", orderID),
|
||||
)
|
||||
|
||||
// Simulate order fill (in real implementation, this would come from IB callbacks)
|
||||
go e.simulateOrderFill(trade)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (e *Executor) simulateOrderFill(trade *model.Trade) {
|
||||
// Wait a bit to simulate order execution
|
||||
time.Sleep(2 * time.Second)
|
||||
|
||||
ctx := context.Background()
|
||||
|
||||
// Update trade as filled
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeFilled
|
||||
trade.FilledAt = &now
|
||||
|
||||
// Simulate executed price (in real implementation, this comes from IB)
|
||||
if trade.TargetPrice != nil {
|
||||
trade.ExecutedPrice = trade.TargetPrice
|
||||
} else {
|
||||
// Use a mock price
|
||||
mockPrice := 150.0
|
||||
trade.ExecutedPrice = &mockPrice
|
||||
}
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
e.logger.Error("failed to update filled trade", slog.Any("error", err))
|
||||
return
|
||||
}
|
||||
|
||||
e.logger.Info("trade filled",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Float64("executed_price", *trade.ExecutedPrice),
|
||||
)
|
||||
|
||||
// Create stop-loss if enabled
|
||||
if trade.StopLossPrice != nil {
|
||||
stopLossPercent := 3.0 // TODO: Get from strategy
|
||||
if err := e.stopLossMgr.CreateStopLoss(ctx, trade, stopLossPercent); err != nil {
|
||||
e.logger.Error("failed to create stop-loss",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
// Mark as completed
|
||||
now = time.Now()
|
||||
trade.Status = model.TradeCompleted
|
||||
trade.CompletedAt = &now
|
||||
|
||||
if err := e.tradeRepo.Update(ctx, trade); err != nil {
|
||||
e.logger.Error("failed to mark trade completed", slog.Any("error", err))
|
||||
return
|
||||
}
|
||||
|
||||
e.logger.Info("trade completed",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
)
|
||||
}
|
||||
|
||||
func (e *Executor) CancelTrade(ctx context.Context, trade *model.Trade) error {
|
||||
if trade.IBOrderID != nil {
|
||||
if err := e.client.CancelOrder(ctx, *trade.IBOrderID); err != nil {
|
||||
return fmt.Errorf("failed to cancel IB order: %w", err)
|
||||
}
|
||||
}
|
||||
|
||||
e.logger.Info("trade cancelled",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,75 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"sync"
|
||||
"time"
|
||||
)
|
||||
|
||||
type RateLimiter struct {
|
||||
maxPerHour int
|
||||
maxParallel int
|
||||
|
||||
mu sync.Mutex
|
||||
hourlyTrades map[int64]int // trades per hour bucket
|
||||
activeTrades int
|
||||
}
|
||||
|
||||
func NewRateLimiter(maxPerHour, maxParallel int) *RateLimiter {
|
||||
return &RateLimiter{
|
||||
maxPerHour: maxPerHour,
|
||||
maxParallel: maxParallel,
|
||||
hourlyTrades: make(map[int64]int),
|
||||
activeTrades: 0,
|
||||
}
|
||||
}
|
||||
|
||||
func (r *RateLimiter) CanTrade() bool {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
// Check parallel limit
|
||||
if r.activeTrades >= r.maxParallel {
|
||||
return false
|
||||
}
|
||||
|
||||
// Check hourly limit
|
||||
currentHour := time.Now().Unix() / 3600
|
||||
if r.hourlyTrades[currentHour] >= r.maxPerHour {
|
||||
return false
|
||||
}
|
||||
|
||||
return true
|
||||
}
|
||||
|
||||
func (r *RateLimiter) RecordTrade() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
currentHour := time.Now().Unix() / 3600
|
||||
r.hourlyTrades[currentHour]++
|
||||
r.activeTrades++
|
||||
|
||||
// Clean up old hour buckets (keep last 2 hours)
|
||||
for hour := range r.hourlyTrades {
|
||||
if hour < currentHour-1 {
|
||||
delete(r.hourlyTrades, hour)
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (r *RateLimiter) ReleaseTrade() {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
if r.activeTrades > 0 {
|
||||
r.activeTrades--
|
||||
}
|
||||
}
|
||||
|
||||
func (r *RateLimiter) GetStats() (hourly int, active int) {
|
||||
r.mu.Lock()
|
||||
defer r.mu.Unlock()
|
||||
|
||||
currentHour := time.Now().Unix() / 3600
|
||||
return r.hourlyTrades[currentHour], r.activeTrades
|
||||
}
|
||||
@@ -0,0 +1,85 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"testing"
|
||||
"time"
|
||||
)
|
||||
|
||||
func TestRateLimiter(t *testing.T) {
|
||||
limiter := NewRateLimiter(3, 2)
|
||||
|
||||
// Should allow first trade
|
||||
if !limiter.CanTrade() {
|
||||
t.Error("Should allow first trade")
|
||||
}
|
||||
|
||||
limiter.RecordTrade()
|
||||
|
||||
// Should allow second trade (within parallel limit)
|
||||
if !limiter.CanTrade() {
|
||||
t.Error("Should allow second trade")
|
||||
}
|
||||
|
||||
limiter.RecordTrade()
|
||||
|
||||
// Should NOT allow third trade (parallel limit reached)
|
||||
if limiter.CanTrade() {
|
||||
t.Error("Should NOT allow third trade (parallel limit)")
|
||||
}
|
||||
|
||||
// Release one trade
|
||||
limiter.ReleaseTrade()
|
||||
|
||||
// Should allow trade again
|
||||
if !limiter.CanTrade() {
|
||||
t.Error("Should allow trade after release")
|
||||
}
|
||||
|
||||
limiter.RecordTrade()
|
||||
|
||||
// Now at hourly limit (3 trades total, 4th trade)
|
||||
limiter.RecordTrade()
|
||||
|
||||
// Release one trade to free up parallel slot
|
||||
limiter.ReleaseTrade()
|
||||
|
||||
// Should NOT allow (hourly limit reached: 4 >= 3)
|
||||
if limiter.CanTrade() {
|
||||
t.Error("Should NOT allow trade (hourly limit)")
|
||||
}
|
||||
|
||||
hourly, active := limiter.GetStats()
|
||||
if hourly != 4 {
|
||||
t.Errorf("Expected 4 hourly trades, got %d", hourly)
|
||||
}
|
||||
if active != 2 {
|
||||
t.Errorf("Expected 2 active trades after one release, got %d", active)
|
||||
}
|
||||
}
|
||||
|
||||
func TestRateLimiterHourBoundary(t *testing.T) {
|
||||
limiter := NewRateLimiter(3, 5)
|
||||
|
||||
// Record 3 trades
|
||||
for i := 0; i < 3; i++ {
|
||||
limiter.RecordTrade()
|
||||
}
|
||||
|
||||
// Should be at limit
|
||||
if limiter.CanTrade() {
|
||||
t.Error("Should be at hourly limit")
|
||||
}
|
||||
|
||||
// Manually manipulate the hour bucket to simulate time passing
|
||||
// In real usage, the hour buckets are cleaned up automatically
|
||||
limiter.mu.Lock()
|
||||
currentHour := time.Now().Unix() / 3600
|
||||
// Clear current hour to simulate new hour
|
||||
delete(limiter.hourlyTrades, currentHour)
|
||||
limiter.mu.Unlock()
|
||||
|
||||
// Should allow trades again in new hour
|
||||
if !limiter.CanTrade() {
|
||||
t.Error("Should allow trades in new hour")
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,127 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/app/client"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type StopLossManager struct {
|
||||
client *client.IBClient
|
||||
tradeRepo *db.TradeRepository
|
||||
logger *slog.Logger
|
||||
enabled bool
|
||||
}
|
||||
|
||||
func NewStopLossManager(client *client.IBClient, tradeRepo *db.TradeRepository, enabled bool, logger *slog.Logger) *StopLossManager {
|
||||
return &StopLossManager{
|
||||
client: client,
|
||||
tradeRepo: tradeRepo,
|
||||
logger: logger,
|
||||
enabled: enabled,
|
||||
}
|
||||
}
|
||||
|
||||
func (s *StopLossManager) CreateStopLoss(ctx context.Context, trade *model.Trade, stopLossPercent float64) error {
|
||||
if !s.enabled || trade.ExecutedPrice == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
executedPrice := *trade.ExecutedPrice
|
||||
stopPrice := executedPrice * (1.0 - stopLossPercent/100.0)
|
||||
|
||||
s.logger.Info("creating stop-loss order",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.Float64("executed_price", executedPrice),
|
||||
slog.Float64("stop_price", stopPrice),
|
||||
)
|
||||
|
||||
order := &client.Order{
|
||||
Symbol: trade.Symbol,
|
||||
Action: model.ActionSell, // Stop-loss is always a sell
|
||||
Quantity: trade.Quantity,
|
||||
OrderType: "STP",
|
||||
StopPrice: &stopPrice,
|
||||
}
|
||||
|
||||
orderID, err := s.client.PlaceOrder(ctx, order)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to place stop-loss order: %w", err)
|
||||
}
|
||||
|
||||
// Store stop-loss details
|
||||
trade.StopLossPrice = &stopPrice
|
||||
trade.IBOrderID = &orderID
|
||||
|
||||
if err := s.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update trade with stop-loss: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *StopLossManager) MonitorStopLoss(ctx context.Context, trade *model.Trade) error {
|
||||
if !s.enabled || trade.StopLossPrice == nil {
|
||||
return nil
|
||||
}
|
||||
|
||||
// In a real implementation, this would check the current market price
|
||||
// and trigger the stop-loss if the price has fallen below the threshold
|
||||
// For now, this is a placeholder
|
||||
|
||||
s.logger.Debug("monitoring stop-loss",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.Float64("stop_price", *trade.StopLossPrice),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (s *StopLossManager) Run(ctx context.Context) error {
|
||||
if !s.enabled {
|
||||
s.logger.Info("stop-loss manager disabled")
|
||||
return nil
|
||||
}
|
||||
|
||||
s.logger.Info("stop-loss manager starting")
|
||||
|
||||
ticker := time.NewTicker(10 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
s.logger.Info("stop-loss manager stopping")
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
if err := s.monitorAllActiveTrades(ctx); err != nil {
|
||||
s.logger.Error("failed to monitor stop-losses", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *StopLossManager) monitorAllActiveTrades(ctx context.Context) error {
|
||||
trades, err := s.tradeRepo.GetActiveTrades(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get active trades: %w", err)
|
||||
}
|
||||
|
||||
for _, trade := range trades {
|
||||
if err := s.MonitorStopLoss(ctx, trade); err != nil {
|
||||
s.logger.Error("failed to monitor stop-loss for trade",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
@@ -0,0 +1,689 @@
|
||||
package trader
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"math/rand"
|
||||
"strings"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/app/client"
|
||||
"github.com/pheinrich/aitrade/pkg/app/strategy"
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type Trader struct {
|
||||
client *client.IBClient
|
||||
tradeRepo *db.TradeRepository
|
||||
newsRepo *db.NewsRepository
|
||||
balanceRepo *db.BalanceRepository
|
||||
whitelistRepo *db.WhitelistRepository
|
||||
positionRepo *db.PositionRepository
|
||||
strategy strategy.Strategy
|
||||
executor *Executor
|
||||
dryRunExec *DryRunExecutor
|
||||
limiter *RateLimiter
|
||||
stopLossMgr *StopLossManager
|
||||
cfg *config.TradingConfig
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func NewTrader(
|
||||
client *client.IBClient,
|
||||
tradeRepo *db.TradeRepository,
|
||||
newsRepo *db.NewsRepository,
|
||||
balanceRepo *db.BalanceRepository,
|
||||
whitelistRepo *db.WhitelistRepository,
|
||||
positionRepo *db.PositionRepository,
|
||||
strategy strategy.Strategy,
|
||||
cfg *config.TradingConfig,
|
||||
logger *slog.Logger,
|
||||
) *Trader {
|
||||
riskParams := strategy.GetRiskParams()
|
||||
limiter := NewRateLimiter(riskParams.MaxTradesPerHour, riskParams.MaxParallelTrades)
|
||||
stopLossMgr := NewStopLossManager(client, tradeRepo, cfg.StopLossEnabled, logger)
|
||||
executor := NewExecutor(client, tradeRepo, stopLossMgr, logger)
|
||||
dryRunExec := NewDryRunExecutor(tradeRepo, balanceRepo, positionRepo, cfg.DryRunBalance, logger)
|
||||
|
||||
return &Trader{
|
||||
client: client,
|
||||
tradeRepo: tradeRepo,
|
||||
newsRepo: newsRepo,
|
||||
balanceRepo: balanceRepo,
|
||||
whitelistRepo: whitelistRepo,
|
||||
positionRepo: positionRepo,
|
||||
strategy: strategy,
|
||||
executor: executor,
|
||||
dryRunExec: dryRunExec,
|
||||
limiter: limiter,
|
||||
stopLossMgr: stopLossMgr,
|
||||
cfg: cfg,
|
||||
logger: logger,
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Trader) Run(ctx context.Context) error {
|
||||
mode := "LIVE"
|
||||
if t.cfg.DryRun {
|
||||
mode = "DRY-RUN"
|
||||
}
|
||||
|
||||
t.logger.Info("trader starting",
|
||||
slog.String("mode", mode),
|
||||
slog.String("strategy", t.strategy.Name()),
|
||||
slog.Duration("pending_time", t.cfg.PendingTime.Duration),
|
||||
slog.Bool("trading_enabled", t.cfg.TradingEnabled),
|
||||
slog.Duration("trading_interval", t.cfg.TradingInterval.Duration),
|
||||
slog.Int("watch_symbols", len(t.cfg.WatchSymbols)),
|
||||
)
|
||||
|
||||
// Wait for IB Gateway connection in live mode (max 10 seconds)
|
||||
if !t.cfg.DryRun {
|
||||
t.logger.Info("waiting for IB Gateway connection")
|
||||
timeout := time.After(10 * time.Second)
|
||||
ticker := time.NewTicker(100 * time.Millisecond)
|
||||
defer ticker.Stop()
|
||||
|
||||
waitLoop:
|
||||
for {
|
||||
if t.client.IsConnected() {
|
||||
t.logger.Info("IB Gateway connected and ready")
|
||||
break waitLoop
|
||||
}
|
||||
|
||||
select {
|
||||
case <-timeout:
|
||||
t.logger.Warn("IB Gateway connection timeout - starting anyway")
|
||||
break waitLoop
|
||||
case <-ctx.Done():
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
// Continue waiting
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
// Start stop-loss manager only in live mode
|
||||
if !t.cfg.DryRun {
|
||||
go t.stopLossMgr.Run(ctx)
|
||||
}
|
||||
|
||||
// Start auto-trading loop if enabled
|
||||
if t.cfg.TradingEnabled {
|
||||
go t.runTradingLoop(ctx)
|
||||
}
|
||||
|
||||
// Process pending trades ticker
|
||||
ticker := time.NewTicker(5 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.logger.Info("trader stopping")
|
||||
return ctx.Err()
|
||||
case <-ticker.C:
|
||||
if err := t.processPendingTrades(ctx); err != nil {
|
||||
t.logger.Error("failed to process pending trades", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Trader) processPendingTrades(ctx context.Context) error {
|
||||
// Get expired pending trades
|
||||
trades, err := t.tradeRepo.GetExpiredPendingTrades(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get expired pending trades: %w", err)
|
||||
}
|
||||
|
||||
for _, trade := range trades {
|
||||
if err := t.executePendingTrade(ctx, trade); err != nil {
|
||||
t.logger.Error("failed to execute pending trade",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) runTradingLoop(ctx context.Context) {
|
||||
t.logger.Info("auto-trading loop starting")
|
||||
|
||||
ticker := time.NewTicker(t.cfg.TradingInterval.Duration)
|
||||
defer ticker.Stop()
|
||||
|
||||
// Run immediately on start
|
||||
if err := t.analyzeAndTrade(ctx); err != nil {
|
||||
t.logger.Error("trading analysis failed", slog.Any("error", err))
|
||||
}
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-ctx.Done():
|
||||
t.logger.Info("auto-trading loop stopping")
|
||||
return
|
||||
case <-ticker.C:
|
||||
if err := t.analyzeAndTrade(ctx); err != nil {
|
||||
t.logger.Error("trading analysis failed", slog.Any("error", err))
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (t *Trader) analyzeAndTrade(ctx context.Context) error {
|
||||
t.logger.Info("analyzing market for trading opportunities")
|
||||
|
||||
// Get recent news (last 24 hours)
|
||||
news, err := t.newsRepo.GetRecent(ctx, 100)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get news: %w", err)
|
||||
}
|
||||
|
||||
t.logger.Info("fetched news articles", slog.Int("count", len(news)))
|
||||
|
||||
// Check existing positions for SELL opportunities
|
||||
if err := t.checkPositionsForSell(ctx, news); err != nil {
|
||||
t.logger.Error("failed to check positions for sell", slog.Any("error", err))
|
||||
}
|
||||
|
||||
// Analyze each watched symbol for BUY opportunities
|
||||
for _, symbol := range t.cfg.WatchSymbols {
|
||||
// Skip if we already have a position
|
||||
position, err := t.positionRepo.GetBySymbol(ctx, symbol)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to check position", slog.String("symbol", symbol), slog.Any("error", err))
|
||||
continue
|
||||
}
|
||||
if position != nil {
|
||||
t.logger.Debug("skipping symbol - position already open", slog.String("symbol", symbol))
|
||||
continue
|
||||
}
|
||||
|
||||
// Filter news relevant to this symbol
|
||||
symbolNews := filterNewsBySymbol(news, symbol)
|
||||
|
||||
t.logger.Info("analyzing symbol",
|
||||
slog.String("symbol", symbol),
|
||||
slog.Int("relevant_news", len(symbolNews)),
|
||||
)
|
||||
|
||||
// Get market data
|
||||
marketData, err := t.getMarketData(ctx, symbol)
|
||||
if err != nil {
|
||||
t.logger.Warn("failed to get market data",
|
||||
slog.String("symbol", symbol),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
// Run strategy analysis
|
||||
signal, err := t.strategy.Analyze(ctx, marketData, symbolNews)
|
||||
if err != nil {
|
||||
t.logger.Warn("strategy analysis failed",
|
||||
slog.String("symbol", symbol),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
// No signal means no trade opportunity
|
||||
if signal == nil {
|
||||
t.logger.Debug("no trade signal generated", slog.String("symbol", symbol))
|
||||
continue
|
||||
}
|
||||
|
||||
// Only process BUY signals here (SELL is handled in checkPositionsForSell)
|
||||
if signal.Action != model.ActionBuy {
|
||||
continue
|
||||
}
|
||||
|
||||
// Calculate position size based on confidence and available capital
|
||||
availableCapital := t.getAvailableCapital(ctx)
|
||||
signal.Quantity = t.strategy.CalculatePositionSize(marketData.LastPrice, signal.Confidence, availableCapital)
|
||||
|
||||
// Apply absolute maximum trade value limit
|
||||
validatedQuantity, err := strategy.ValidateTradeValue(signal.Quantity, marketData.LastPrice, t.cfg.MaxTradeValue)
|
||||
if err != nil {
|
||||
t.logger.Warn("trade rejected - exceeds max trade value",
|
||||
slog.String("symbol", signal.Symbol),
|
||||
slog.Int("calculated_quantity", signal.Quantity),
|
||||
slog.Float64("price", marketData.LastPrice),
|
||||
slog.Float64("max_trade_value", t.cfg.MaxTradeValue),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
signal.Quantity = validatedQuantity
|
||||
tradeValue := marketData.LastPrice * float64(signal.Quantity)
|
||||
|
||||
// Create pending trade
|
||||
t.logger.Info("trade signal generated",
|
||||
slog.String("symbol", signal.Symbol),
|
||||
slog.String("action", string(signal.Action)),
|
||||
slog.Int("quantity", signal.Quantity),
|
||||
slog.Float64("confidence", signal.Confidence),
|
||||
slog.Float64("trade_value", tradeValue),
|
||||
)
|
||||
|
||||
if err := t.CreatePendingTrade(ctx, signal); err != nil {
|
||||
t.logger.Error("failed to create pending trade",
|
||||
slog.String("symbol", signal.Symbol),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) checkPositionsForSell(ctx context.Context, news []*model.NewsArticle) error {
|
||||
positions, err := t.positionRepo.GetAll(ctx)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get positions: %w", err)
|
||||
}
|
||||
|
||||
for _, position := range positions {
|
||||
// Get current market data
|
||||
marketData, err := t.getMarketData(ctx, position.Symbol)
|
||||
if err != nil {
|
||||
t.logger.Warn("failed to get market data for position",
|
||||
slog.String("symbol", position.Symbol),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
// Update position with current price
|
||||
position.CurrentPrice = &marketData.LastPrice
|
||||
pnl := (marketData.LastPrice - position.EntryPrice) * float64(position.Quantity)
|
||||
position.UnrealizedPnL = &pnl
|
||||
|
||||
if err := t.positionRepo.Update(ctx, position); err != nil {
|
||||
t.logger.Error("failed to update position", slog.Any("error", err))
|
||||
}
|
||||
|
||||
// Check hold time
|
||||
holdTime := time.Since(position.OpenedAt)
|
||||
if holdTime < time.Duration(t.cfg.HoldTimeMinutes)*time.Minute {
|
||||
t.logger.Debug("position not old enough to sell",
|
||||
slog.String("symbol", position.Symbol),
|
||||
slog.Duration("hold_time", holdTime),
|
||||
)
|
||||
continue
|
||||
}
|
||||
|
||||
// Calculate profit percentage
|
||||
profitPct := ((marketData.LastPrice - position.EntryPrice) / position.EntryPrice) * 100
|
||||
|
||||
// SELL trigger 1: Take profit target reached
|
||||
if profitPct >= t.cfg.TakeProfitPercent {
|
||||
t.logger.Info("take profit triggered",
|
||||
slog.String("symbol", position.Symbol),
|
||||
slog.Float64("profit_pct", profitPct),
|
||||
slog.Float64("target_pct", t.cfg.TakeProfitPercent),
|
||||
)
|
||||
|
||||
signal := &strategy.TradeSignal{
|
||||
Symbol: position.Symbol,
|
||||
Action: model.ActionSell,
|
||||
Quantity: position.Quantity,
|
||||
Confidence: 0.95,
|
||||
Reasoning: fmt.Sprintf("Take profit: %.2f%% gain (target: %.2f%%)", profitPct, t.cfg.TakeProfitPercent),
|
||||
}
|
||||
|
||||
if err := t.CreatePendingTrade(ctx, signal); err != nil {
|
||||
t.logger.Error("failed to create sell trade", slog.Any("error", err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
|
||||
// SELL trigger 2: Negative sentiment (if enabled)
|
||||
if t.cfg.NegativeSentiment {
|
||||
symbolNews := filterNewsBySymbol(news, position.Symbol)
|
||||
if len(symbolNews) >= 2 {
|
||||
signal, err := t.strategy.Analyze(ctx, marketData, symbolNews)
|
||||
if err == nil && signal != nil && signal.Action == model.ActionSell {
|
||||
t.logger.Info("negative sentiment sell triggered",
|
||||
slog.String("symbol", position.Symbol),
|
||||
slog.String("reasoning", signal.Reasoning),
|
||||
)
|
||||
|
||||
signal.Quantity = position.Quantity
|
||||
if err := t.CreatePendingTrade(ctx, signal); err != nil {
|
||||
t.logger.Error("failed to create sell trade", slog.Any("error", err))
|
||||
}
|
||||
continue
|
||||
}
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) getAvailableCapital(ctx context.Context) float64 {
|
||||
var totalCapital float64
|
||||
|
||||
if t.cfg.DryRun {
|
||||
totalCapital = t.dryRunExec.GetBalance()
|
||||
} else {
|
||||
// In live mode, get from balance repository
|
||||
balance, err := t.balanceRepo.GetLatest(ctx)
|
||||
if err != nil || balance == nil {
|
||||
t.logger.Warn("failed to get balance, using default", slog.Any("error", err))
|
||||
totalCapital = 100000.0 // Fallback
|
||||
} else {
|
||||
// Use buying power (includes margin) for available capital
|
||||
totalCapital = balance.BuyingPower
|
||||
}
|
||||
}
|
||||
|
||||
// Get current open positions count
|
||||
positions, err := t.positionRepo.GetAll(ctx)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to get positions", slog.Any("error", err))
|
||||
positions = []*model.Position{} // Empty on error
|
||||
}
|
||||
openPositions := len(positions)
|
||||
|
||||
// Get pending trades count
|
||||
pendingTrades, err := t.tradeRepo.GetPendingTrades(ctx)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to get pending trades", slog.Any("error", err))
|
||||
pendingTrades = []*model.Trade{} // Empty on error
|
||||
}
|
||||
pendingCount := len(pendingTrades)
|
||||
|
||||
// Calculate slots: max parallel trades - (open positions + pending trades)
|
||||
maxParallel := t.limiter.maxParallel
|
||||
usedSlots := openPositions + pendingCount
|
||||
availableSlots := maxParallel - usedSlots
|
||||
|
||||
if availableSlots <= 0 {
|
||||
t.logger.Warn("no available trade slots",
|
||||
slog.Int("max_parallel", maxParallel),
|
||||
slog.Int("open_positions", openPositions),
|
||||
slog.Int("pending_trades", pendingCount),
|
||||
)
|
||||
return 0.0 // No capital available if no slots
|
||||
}
|
||||
|
||||
// Divide total capital by max parallel trades to reserve capital for other trades
|
||||
// This ensures we don't use all capital on the first trade
|
||||
capitalPerSlot := totalCapital / float64(maxParallel)
|
||||
|
||||
t.logger.Debug("calculated available capital",
|
||||
slog.Float64("total_capital", totalCapital),
|
||||
slog.Int("max_parallel", maxParallel),
|
||||
slog.Int("available_slots", availableSlots),
|
||||
slog.Float64("capital_per_slot", capitalPerSlot),
|
||||
)
|
||||
|
||||
return capitalPerSlot
|
||||
}
|
||||
|
||||
func (t *Trader) getMarketData(ctx context.Context, symbol string) (*strategy.MarketData, error) {
|
||||
// In dry-run mode, generate simulated market data
|
||||
if t.cfg.DryRun {
|
||||
return t.generateSimulatedMarketData(symbol), nil
|
||||
}
|
||||
|
||||
// In live mode, get real market data from IB Gateway
|
||||
ibData, err := t.client.GetMarketData(ctx, symbol)
|
||||
if err != nil {
|
||||
t.logger.Warn("failed to get real market data from IB, using simulated data",
|
||||
slog.String("symbol", symbol),
|
||||
slog.Any("error", err))
|
||||
return t.generateSimulatedMarketData(symbol), nil
|
||||
}
|
||||
|
||||
// Convert IB market data to strategy market data format
|
||||
return &strategy.MarketData{
|
||||
Symbol: ibData.Symbol,
|
||||
LastPrice: ibData.LastPrice,
|
||||
BidPrice: ibData.BidPrice,
|
||||
AskPrice: ibData.AskPrice,
|
||||
Volume: ibData.Volume,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (t *Trader) generateSimulatedMarketData(symbol string) *strategy.MarketData {
|
||||
// Generate realistic simulated prices based on symbol
|
||||
basePrice := 150.0
|
||||
switch symbol {
|
||||
case "AAPL":
|
||||
basePrice = 180.0
|
||||
case "MSFT":
|
||||
basePrice = 420.0
|
||||
case "GOOGL":
|
||||
basePrice = 175.0
|
||||
case "TSLA":
|
||||
basePrice = 250.0
|
||||
case "AMZN":
|
||||
basePrice = 190.0
|
||||
}
|
||||
|
||||
// Add some randomness
|
||||
change := (rand.Float64() - 0.5) * 10.0
|
||||
lastPrice := basePrice + change
|
||||
changePercent := (change / basePrice) * 100.0
|
||||
|
||||
return &strategy.MarketData{
|
||||
Symbol: symbol,
|
||||
LastPrice: lastPrice,
|
||||
BidPrice: lastPrice - 0.5,
|
||||
AskPrice: lastPrice + 0.5,
|
||||
Volume: int64(rand.Intn(10000000) + 1000000),
|
||||
Change: change,
|
||||
ChangePercent: changePercent,
|
||||
}
|
||||
}
|
||||
|
||||
func filterNewsBySymbol(news []*model.NewsArticle, symbol string) []*model.NewsArticle {
|
||||
var filtered []*model.NewsArticle
|
||||
for _, article := range news {
|
||||
if article.Symbols != "" && strings.Contains(article.Symbols, symbol) {
|
||||
filtered = append(filtered, article)
|
||||
}
|
||||
}
|
||||
return filtered
|
||||
}
|
||||
|
||||
func (t *Trader) executePendingTrade(ctx context.Context, trade *model.Trade) error {
|
||||
// Check whitelist before execution
|
||||
whitelisted, err := t.whitelistRepo.IsSymbolWhitelisted(ctx, trade.Symbol)
|
||||
if err != nil {
|
||||
t.logger.Error("failed to check whitelist",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
slog.Any("error", err),
|
||||
)
|
||||
return fmt.Errorf("failed to check whitelist: %w", err)
|
||||
}
|
||||
|
||||
if !whitelisted {
|
||||
t.logger.Warn("trade rejected - symbol not whitelisted",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.String("symbol", trade.Symbol),
|
||||
)
|
||||
|
||||
// Mark trade as rejected
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeRejected
|
||||
trade.RejectedAt = &now
|
||||
reason := fmt.Sprintf("Symbol %s not in whitelist", trade.Symbol)
|
||||
trade.RejectionReason = &reason
|
||||
|
||||
if err := t.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update rejected trade: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// Check rate limits
|
||||
if !t.limiter.CanTrade() {
|
||||
hourly, active := t.limiter.GetStats()
|
||||
t.logger.Warn("rate limit exceeded, postponing trade",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Int("hourly_trades", hourly),
|
||||
slog.Int("active_trades", active),
|
||||
)
|
||||
|
||||
// Extend pending time by 1 minute
|
||||
newPendingUntil := time.Now().Add(1 * time.Minute)
|
||||
trade.PendingUtil = &newPendingUntil
|
||||
return t.tradeRepo.Update(ctx, trade)
|
||||
}
|
||||
|
||||
// Execute the trade (dry run or live)
|
||||
t.limiter.RecordTrade()
|
||||
|
||||
if t.cfg.DryRun {
|
||||
err = t.dryRunExec.ExecuteTrade(ctx, trade)
|
||||
} else {
|
||||
err = t.executor.ExecuteTrade(ctx, trade)
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
t.limiter.ReleaseTrade()
|
||||
return fmt.Errorf("failed to execute trade: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) CreatePendingTrade(ctx context.Context, signal *strategy.TradeSignal) error {
|
||||
mode := "LIVE"
|
||||
if t.cfg.DryRun {
|
||||
mode = "DRY-RUN"
|
||||
}
|
||||
|
||||
t.logger.Info("creating pending trade",
|
||||
slog.String("mode", mode),
|
||||
slog.String("symbol", signal.Symbol),
|
||||
slog.String("action", string(signal.Action)),
|
||||
slog.Int("quantity", signal.Quantity),
|
||||
slog.Float64("confidence", signal.Confidence),
|
||||
)
|
||||
|
||||
now := time.Now()
|
||||
pendingUntil := now.Add(t.cfg.PendingTime.Duration)
|
||||
|
||||
trade := &model.Trade{
|
||||
Symbol: signal.Symbol,
|
||||
Action: signal.Action,
|
||||
Quantity: signal.Quantity,
|
||||
Status: model.TradePending,
|
||||
Confidence: signal.Confidence,
|
||||
Reasoning: signal.Reasoning,
|
||||
CreatedAt: now,
|
||||
PendingUtil: &pendingUntil,
|
||||
IsDryRun: t.cfg.DryRun,
|
||||
}
|
||||
|
||||
// Set stop-loss price if enabled
|
||||
if t.strategy.GetRiskParams().StopLossPercent > 0 {
|
||||
// This will be calculated when we have the executed price
|
||||
stopLossPercent := t.strategy.GetRiskParams().StopLossPercent
|
||||
_ = stopLossPercent // Will be used after execution
|
||||
}
|
||||
|
||||
if err := t.tradeRepo.Create(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to create pending trade: %w", err)
|
||||
}
|
||||
|
||||
t.logger.Info("pending trade created",
|
||||
slog.Int64("trade_id", trade.ID),
|
||||
slog.Time("pending_until", pendingUntil),
|
||||
slog.Bool("dry_run", t.cfg.DryRun),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) ApproveTrade(ctx context.Context, tradeID int64, forceNow bool) error {
|
||||
trade, err := t.tradeRepo.GetByID(ctx, tradeID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get trade: %w", err)
|
||||
}
|
||||
|
||||
if trade.Status != model.TradePending {
|
||||
return fmt.Errorf("trade %d is not in pending status", tradeID)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeApproved
|
||||
trade.ApprovedAt = &now
|
||||
trade.ForcedByUser = forceNow
|
||||
|
||||
if forceNow {
|
||||
// Execute immediately
|
||||
trade.PendingUtil = &now
|
||||
}
|
||||
|
||||
if err := t.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update trade: %w", err)
|
||||
}
|
||||
|
||||
t.logger.Info("trade approved",
|
||||
slog.Int64("trade_id", tradeID),
|
||||
slog.Bool("force_now", forceNow),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (t *Trader) RejectTrade(ctx context.Context, tradeID int64, reason string) error {
|
||||
trade, err := t.tradeRepo.GetByID(ctx, tradeID)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get trade: %w", err)
|
||||
}
|
||||
|
||||
if trade.Status != model.TradePending {
|
||||
return fmt.Errorf("trade %d is not in pending status", tradeID)
|
||||
}
|
||||
|
||||
now := time.Now()
|
||||
trade.Status = model.TradeRejected
|
||||
trade.RejectedAt = &now
|
||||
trade.RejectionReason = &reason
|
||||
|
||||
if err := t.tradeRepo.Update(ctx, trade); err != nil {
|
||||
return fmt.Errorf("failed to update trade: %w", err)
|
||||
}
|
||||
|
||||
t.logger.Info("trade rejected",
|
||||
slog.Int64("trade_id", tradeID),
|
||||
slog.String("reason", reason),
|
||||
)
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
// GetLiveBalance fetches the current account balance from IB Gateway
|
||||
func (t *Trader) GetLiveBalance(ctx context.Context) (*model.Balance, error) {
|
||||
summary, err := t.client.GetAccountSummary(ctx)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get account summary from IB: %w", err)
|
||||
}
|
||||
|
||||
balance := &model.Balance{
|
||||
Timestamp: time.Now(),
|
||||
TotalValue: summary.TotalValue,
|
||||
CashBalance: summary.CashBalance,
|
||||
BuyingPower: summary.BuyingPower,
|
||||
UnrealizedPnL: &summary.UnrealizedPnL,
|
||||
RealizedPnL: &summary.RealizedPnL,
|
||||
}
|
||||
|
||||
return balance, nil
|
||||
}
|
||||
@@ -0,0 +1,189 @@
|
||||
package web
|
||||
|
||||
import (
|
||||
"context"
|
||||
"crypto/rand"
|
||||
"encoding/base64"
|
||||
"fmt"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"time"
|
||||
|
||||
"github.com/coreos/go-oidc/v3/oidc"
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"golang.org/x/oauth2"
|
||||
)
|
||||
|
||||
type AuthMiddleware struct {
|
||||
enabled bool
|
||||
provider *oidc.Provider
|
||||
oauth2Config oauth2.Config
|
||||
verifier *oidc.IDTokenVerifier
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func NewAuthMiddleware(cfg *config.OIDCConfig, logger *slog.Logger) (*AuthMiddleware, error) {
|
||||
if !cfg.Enabled {
|
||||
logger.Info("OIDC authentication disabled")
|
||||
return &AuthMiddleware{
|
||||
enabled: false,
|
||||
logger: logger,
|
||||
}, nil
|
||||
}
|
||||
|
||||
ctx := context.Background()
|
||||
provider, err := oidc.NewProvider(ctx, cfg.Issuer)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create OIDC provider: %w", err)
|
||||
}
|
||||
|
||||
oauth2Config := oauth2.Config{
|
||||
ClientID: cfg.ClientID,
|
||||
ClientSecret: cfg.ClientSecret,
|
||||
RedirectURL: cfg.RedirectURL,
|
||||
Endpoint: provider.Endpoint(),
|
||||
Scopes: cfg.Scopes,
|
||||
}
|
||||
|
||||
verifier := provider.Verifier(&oidc.Config{ClientID: cfg.ClientID})
|
||||
|
||||
logger.Info("OIDC authentication enabled",
|
||||
slog.String("issuer", cfg.Issuer),
|
||||
slog.String("client_id", cfg.ClientID),
|
||||
)
|
||||
|
||||
return &AuthMiddleware{
|
||||
enabled: true,
|
||||
provider: provider,
|
||||
oauth2Config: oauth2Config,
|
||||
verifier: verifier,
|
||||
logger: logger,
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (a *AuthMiddleware) Middleware(next http.Handler) http.Handler {
|
||||
return http.HandlerFunc(func(w http.ResponseWriter, r *http.Request) {
|
||||
// If OIDC is disabled, allow all requests
|
||||
if !a.enabled {
|
||||
next.ServeHTTP(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// Check for session cookie
|
||||
cookie, err := r.Cookie("aitrade_session")
|
||||
if err != nil {
|
||||
// No session, redirect to login
|
||||
a.redirectToLogin(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// Verify the ID token
|
||||
ctx := r.Context()
|
||||
_, err = a.verifier.Verify(ctx, cookie.Value)
|
||||
if err != nil {
|
||||
a.logger.Warn("invalid session token", slog.Any("error", err))
|
||||
a.redirectToLogin(w, r)
|
||||
return
|
||||
}
|
||||
|
||||
// Token is valid, proceed
|
||||
next.ServeHTTP(w, r)
|
||||
})
|
||||
}
|
||||
|
||||
func (a *AuthMiddleware) redirectToLogin(w http.ResponseWriter, r *http.Request) {
|
||||
state, err := generateRandomState()
|
||||
if err != nil {
|
||||
http.Error(w, "Failed to generate state", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Store state in cookie for verification
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "oidc_state",
|
||||
Value: state,
|
||||
Path: "/",
|
||||
MaxAge: 300, // 5 minutes
|
||||
HttpOnly: true,
|
||||
Secure: false, // Set to true in production with HTTPS
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
})
|
||||
|
||||
authURL := a.oauth2Config.AuthCodeURL(state)
|
||||
http.Redirect(w, r, authURL, http.StatusFound)
|
||||
}
|
||||
|
||||
func (a *AuthMiddleware) HandleCallback(w http.ResponseWriter, r *http.Request) {
|
||||
if !a.enabled {
|
||||
http.Error(w, "OIDC not enabled", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Verify state
|
||||
stateCookie, err := r.Cookie("oidc_state")
|
||||
if err != nil {
|
||||
http.Error(w, "State cookie not found", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if r.URL.Query().Get("state") != stateCookie.Value {
|
||||
http.Error(w, "State mismatch", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
// Exchange code for token
|
||||
ctx := r.Context()
|
||||
oauth2Token, err := a.oauth2Config.Exchange(ctx, r.URL.Query().Get("code"))
|
||||
if err != nil {
|
||||
a.logger.Error("failed to exchange code", slog.Any("error", err))
|
||||
http.Error(w, "Failed to exchange code", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Extract ID token
|
||||
rawIDToken, ok := oauth2Token.Extra("id_token").(string)
|
||||
if !ok {
|
||||
http.Error(w, "No id_token in response", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Verify ID token
|
||||
idToken, err := a.verifier.Verify(ctx, rawIDToken)
|
||||
if err != nil {
|
||||
a.logger.Error("failed to verify ID token", slog.Any("error", err))
|
||||
http.Error(w, "Failed to verify ID token", http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
// Store token in session cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "aitrade_session",
|
||||
Value: rawIDToken,
|
||||
Path: "/",
|
||||
MaxAge: int(time.Until(idToken.Expiry).Seconds()),
|
||||
HttpOnly: true,
|
||||
Secure: false, // Set to true in production with HTTPS
|
||||
SameSite: http.SameSiteLaxMode,
|
||||
})
|
||||
|
||||
// Clear state cookie
|
||||
http.SetCookie(w, &http.Cookie{
|
||||
Name: "oidc_state",
|
||||
Value: "",
|
||||
Path: "/",
|
||||
MaxAge: -1,
|
||||
})
|
||||
|
||||
a.logger.Info("user authenticated successfully")
|
||||
|
||||
// Redirect to home
|
||||
http.Redirect(w, r, "/", http.StatusFound)
|
||||
}
|
||||
|
||||
func generateRandomState() (string, error) {
|
||||
b := make([]byte, 32)
|
||||
if _, err := rand.Read(b); err != nil {
|
||||
return "", err
|
||||
}
|
||||
return base64.URLEncoding.EncodeToString(b), nil
|
||||
}
|
||||
@@ -0,0 +1,496 @@
|
||||
package web
|
||||
|
||||
import (
|
||||
"context"
|
||||
"embed"
|
||||
"encoding/json"
|
||||
"fmt"
|
||||
"html/template"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"net/http"
|
||||
"sync"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/app/trader"
|
||||
"github.com/pheinrich/aitrade/pkg/config"
|
||||
"github.com/pheinrich/aitrade/pkg/db"
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
//go:embed templates/*.html static/*
|
||||
var embeddedFS embed.FS
|
||||
|
||||
type Server struct {
|
||||
cfg *config.WebConfig
|
||||
tradingCfg *config.TradingConfig
|
||||
auth *AuthMiddleware
|
||||
tradeRepo *db.TradeRepository
|
||||
balanceRepo *db.BalanceRepository
|
||||
newsRepo *db.NewsRepository
|
||||
whitelistRepo *db.WhitelistRepository
|
||||
trader *trader.Trader
|
||||
logger *slog.Logger
|
||||
templates *template.Template
|
||||
|
||||
// SSE
|
||||
sseClients map[chan string]bool
|
||||
sseClientsMu sync.Mutex
|
||||
}
|
||||
|
||||
func NewServer(
|
||||
cfg *config.WebConfig,
|
||||
tradingCfg *config.TradingConfig,
|
||||
oidcCfg *config.OIDCConfig,
|
||||
tradeRepo *db.TradeRepository,
|
||||
balanceRepo *db.BalanceRepository,
|
||||
newsRepo *db.NewsRepository,
|
||||
whitelistRepo *db.WhitelistRepository,
|
||||
trader *trader.Trader,
|
||||
logger *slog.Logger,
|
||||
) (*Server, error) {
|
||||
auth, err := NewAuthMiddleware(oidcCfg, logger)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to create auth middleware: %w", err)
|
||||
}
|
||||
|
||||
// Parse templates
|
||||
tmpl, err := template.ParseFS(embeddedFS, "templates/*.html")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to parse templates: %w", err)
|
||||
}
|
||||
|
||||
return &Server{
|
||||
cfg: cfg,
|
||||
tradingCfg: tradingCfg,
|
||||
auth: auth,
|
||||
tradeRepo: tradeRepo,
|
||||
balanceRepo: balanceRepo,
|
||||
newsRepo: newsRepo,
|
||||
whitelistRepo: whitelistRepo,
|
||||
trader: trader,
|
||||
logger: logger,
|
||||
templates: tmpl,
|
||||
sseClients: make(map[chan string]bool),
|
||||
}, nil
|
||||
}
|
||||
|
||||
func (s *Server) Run(ctx context.Context) error {
|
||||
mux := http.NewServeMux()
|
||||
|
||||
// Static files
|
||||
staticFS, err := fs.Sub(embeddedFS, "static")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get static filesystem: %w", err)
|
||||
}
|
||||
mux.Handle("/static/", http.StripPrefix("/static/", http.FileServer(http.FS(staticFS))))
|
||||
|
||||
// Public endpoints (no auth required)
|
||||
mux.HandleFunc("/health", s.handleHealth)
|
||||
mux.HandleFunc("/callback", s.auth.HandleCallback)
|
||||
|
||||
// Protected endpoints (auth required)
|
||||
mux.Handle("/", s.auth.Middleware(http.HandlerFunc(s.handleOverview)))
|
||||
mux.Handle("/overview", s.auth.Middleware(http.HandlerFunc(s.handleOverview)))
|
||||
mux.Handle("/trades", s.auth.Middleware(http.HandlerFunc(s.handleTradesPage)))
|
||||
mux.Handle("/whitelist", s.auth.Middleware(http.HandlerFunc(s.handleWhitelistPage)))
|
||||
mux.Handle("/news", s.auth.Middleware(http.HandlerFunc(s.handleNewsPage)))
|
||||
|
||||
// Content endpoints (partial HTML for SPA)
|
||||
mux.Handle("/content/overview", s.auth.Middleware(http.HandlerFunc(s.handleOverviewContent)))
|
||||
mux.Handle("/content/trades", s.auth.Middleware(http.HandlerFunc(s.handleTradesContent)))
|
||||
mux.Handle("/content/whitelist", s.auth.Middleware(http.HandlerFunc(s.handleWhitelistContent)))
|
||||
mux.Handle("/content/news", s.auth.Middleware(http.HandlerFunc(s.handleNewsContent)))
|
||||
|
||||
mux.Handle("/events", s.auth.Middleware(http.HandlerFunc(s.handleSSE)))
|
||||
mux.Handle("/api/balance", s.auth.Middleware(http.HandlerFunc(s.handleGetBalance)))
|
||||
mux.Handle("/api/config", s.auth.Middleware(http.HandlerFunc(s.handleGetConfig)))
|
||||
mux.Handle("/api/news", s.auth.Middleware(http.HandlerFunc(s.handleGetNews)))
|
||||
mux.Handle("/api/news/recent", s.auth.Middleware(http.HandlerFunc(s.handleGetNewsRecent)))
|
||||
mux.Handle("/api/trades", s.auth.Middleware(http.HandlerFunc(s.handleGetTrades)))
|
||||
mux.Handle("/api/trades/", s.auth.Middleware(http.HandlerFunc(s.handleTradeAction)))
|
||||
mux.Handle("/api/whitelist", s.auth.Middleware(http.HandlerFunc(s.handleWhitelist)))
|
||||
mux.Handle("/api/whitelist/", s.auth.Middleware(http.HandlerFunc(s.handleWhitelistItem)))
|
||||
|
||||
addr := fmt.Sprintf("%s:%s", s.cfg.Host, s.cfg.Port)
|
||||
server := &http.Server{
|
||||
Addr: addr,
|
||||
Handler: mux,
|
||||
}
|
||||
|
||||
s.logger.Info("web server starting", slog.String("addr", addr))
|
||||
|
||||
// Start server in goroutine
|
||||
go func() {
|
||||
if err := server.ListenAndServe(); err != nil && err != http.ErrServerClosed {
|
||||
s.logger.Error("web server error", slog.Any("error", err))
|
||||
}
|
||||
}()
|
||||
|
||||
// Wait for context cancellation
|
||||
<-ctx.Done()
|
||||
|
||||
s.logger.Info("web server stopping")
|
||||
|
||||
// Shutdown gracefully
|
||||
shutdownCtx, cancel := context.WithTimeout(context.Background(), 5*time.Second)
|
||||
defer cancel()
|
||||
|
||||
return server.Shutdown(shutdownCtx)
|
||||
}
|
||||
|
||||
func (s *Server) handleHealth(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]string{
|
||||
"status": "healthy",
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Server) handleOverview(w http.ResponseWriter, r *http.Request) {
|
||||
data := map[string]interface{}{
|
||||
"Title": "Overview",
|
||||
"CurrentPage": "overview",
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "shell.html", data); err != nil {
|
||||
s.logger.Error("failed to render template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleTradesPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := map[string]interface{}{
|
||||
"Title": "Trades",
|
||||
"CurrentPage": "trades",
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "shell.html", data); err != nil {
|
||||
s.logger.Error("failed to render template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleWhitelistPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := map[string]interface{}{
|
||||
"Title": "Whitelist",
|
||||
"CurrentPage": "whitelist",
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "shell.html", data); err != nil {
|
||||
s.logger.Error("failed to render template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleNewsPage(w http.ResponseWriter, r *http.Request) {
|
||||
data := map[string]interface{}{
|
||||
"Title": "News Feed",
|
||||
"CurrentPage": "news",
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "shell.html", data); err != nil {
|
||||
s.logger.Error("failed to render template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
// Content endpoints (return partial HTML for SPA)
|
||||
func (s *Server) handleOverviewContent(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "overview-content", nil); err != nil {
|
||||
s.logger.Error("failed to render content template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleTradesContent(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "trades-content", nil); err != nil {
|
||||
s.logger.Error("failed to render content template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleWhitelistContent(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "whitelist-content", nil); err != nil {
|
||||
s.logger.Error("failed to render content template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleNewsContent(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/html; charset=utf-8")
|
||||
if err := s.templates.ExecuteTemplate(w, "news-content", nil); err != nil {
|
||||
s.logger.Error("failed to render content template", slog.Any("error", err))
|
||||
http.Error(w, "Internal Server Error", http.StatusInternalServerError)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleGetTrades(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Get trades from last 7 days
|
||||
since := time.Now().AddDate(0, 0, -7)
|
||||
trades, err := s.tradeRepo.GetTradesSince(ctx, since)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(trades)
|
||||
}
|
||||
|
||||
func (s *Server) handleGetBalance(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Get live balance from IB Gateway via trader
|
||||
balance, err := s.trader.GetLiveBalance(ctx)
|
||||
if err != nil {
|
||||
s.logger.Warn("failed to get live balance from IB, falling back to DB",
|
||||
slog.Any("error", err))
|
||||
|
||||
// Fallback: read from DB if IB Gateway is unavailable
|
||||
balance, err = s.balanceRepo.GetLatest(ctx)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
if balance == nil {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]float64{"TotalValue": 0.0})
|
||||
return
|
||||
}
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(balance)
|
||||
}
|
||||
|
||||
func (s *Server) handleGetConfig(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(map[string]interface{}{
|
||||
"dry_run": s.tradingCfg.DryRun,
|
||||
})
|
||||
}
|
||||
|
||||
func (s *Server) handleGetNews(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
news, err := s.newsRepo.GetRecent(ctx, 20)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(news)
|
||||
}
|
||||
|
||||
func (s *Server) handleGetNewsRecent(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Get limit from query parameter, default to 50
|
||||
limit := 50
|
||||
if limitStr := r.URL.Query().Get("limit"); limitStr != "" {
|
||||
fmt.Sscanf(limitStr, "%d", &limit)
|
||||
}
|
||||
|
||||
news, err := s.newsRepo.GetRecent(ctx, limit)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(news)
|
||||
}
|
||||
|
||||
func (s *Server) handleTradeAction(w http.ResponseWriter, r *http.Request) {
|
||||
// Parse trade ID from URL: /api/trades/{id}/{action}
|
||||
path := r.URL.Path[len("/api/trades/"):]
|
||||
|
||||
var tradeID int64
|
||||
var action string
|
||||
fmt.Sscanf(path, "%d/%s", &tradeID, &action)
|
||||
|
||||
ctx := r.Context()
|
||||
|
||||
switch action {
|
||||
case "approve":
|
||||
forceNow := r.URL.Query().Get("force") == "true"
|
||||
if err := s.trader.ApproveTrade(ctx, tradeID, forceNow); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
s.BroadcastSSE("trade.approved")
|
||||
|
||||
case "reject":
|
||||
var body struct {
|
||||
Reason string `json:"reason"`
|
||||
}
|
||||
json.NewDecoder(r.Body).Decode(&body)
|
||||
|
||||
if err := s.trader.RejectTrade(ctx, tradeID, body.Reason); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
s.BroadcastSSE("trade.rejected")
|
||||
|
||||
default:
|
||||
http.Error(w, "Invalid action", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
w.WriteHeader(http.StatusOK)
|
||||
}
|
||||
|
||||
func (s *Server) handleWhitelist(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
// Get all whitelist entries
|
||||
entries, err := s.whitelistRepo.GetAll(ctx)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(entries)
|
||||
|
||||
case http.MethodPost:
|
||||
// Create new whitelist entry
|
||||
var entry model.WhitelistEntry
|
||||
if err := json.NewDecoder(r.Body).Decode(&entry); err != nil {
|
||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
if err := s.whitelistRepo.Create(ctx, &entry); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
s.BroadcastSSE("whitelist.updated")
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(entry)
|
||||
|
||||
default:
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleWhitelistItem(w http.ResponseWriter, r *http.Request) {
|
||||
ctx := r.Context()
|
||||
|
||||
// Parse ID from URL: /api/whitelist/{id}
|
||||
path := r.URL.Path[len("/api/whitelist/"):]
|
||||
var id int64
|
||||
fmt.Sscanf(path, "%d", &id)
|
||||
|
||||
switch r.Method {
|
||||
case http.MethodGet:
|
||||
// Get single whitelist entry
|
||||
entry, err := s.whitelistRepo.GetByID(ctx, id)
|
||||
if err != nil {
|
||||
http.Error(w, err.Error(), http.StatusNotFound)
|
||||
return
|
||||
}
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(entry)
|
||||
|
||||
case http.MethodPut:
|
||||
// Update whitelist entry
|
||||
var entry model.WhitelistEntry
|
||||
if err := json.NewDecoder(r.Body).Decode(&entry); err != nil {
|
||||
http.Error(w, "Invalid request body", http.StatusBadRequest)
|
||||
return
|
||||
}
|
||||
|
||||
entry.ID = id
|
||||
if err := s.whitelistRepo.Update(ctx, &entry); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
s.BroadcastSSE("whitelist.updated")
|
||||
|
||||
w.Header().Set("Content-Type", "application/json")
|
||||
json.NewEncoder(w).Encode(entry)
|
||||
|
||||
case http.MethodDelete:
|
||||
// Delete whitelist entry
|
||||
if err := s.whitelistRepo.Delete(ctx, id); err != nil {
|
||||
http.Error(w, err.Error(), http.StatusInternalServerError)
|
||||
return
|
||||
}
|
||||
|
||||
s.BroadcastSSE("whitelist.updated")
|
||||
|
||||
w.WriteHeader(http.StatusNoContent)
|
||||
|
||||
default:
|
||||
http.Error(w, "Method not allowed", http.StatusMethodNotAllowed)
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) handleSSE(w http.ResponseWriter, r *http.Request) {
|
||||
w.Header().Set("Content-Type", "text/event-stream")
|
||||
w.Header().Set("Cache-Control", "no-cache")
|
||||
w.Header().Set("Connection", "keep-alive")
|
||||
|
||||
// Create client channel
|
||||
clientChan := make(chan string)
|
||||
|
||||
s.sseClientsMu.Lock()
|
||||
s.sseClients[clientChan] = true
|
||||
s.sseClientsMu.Unlock()
|
||||
|
||||
// Remove client on disconnect
|
||||
defer func() {
|
||||
s.sseClientsMu.Lock()
|
||||
delete(s.sseClients, clientChan)
|
||||
s.sseClientsMu.Unlock()
|
||||
close(clientChan)
|
||||
}()
|
||||
|
||||
// Send keep-alive
|
||||
ticker := time.NewTicker(30 * time.Second)
|
||||
defer ticker.Stop()
|
||||
|
||||
for {
|
||||
select {
|
||||
case <-r.Context().Done():
|
||||
return
|
||||
case msg := <-clientChan:
|
||||
fmt.Fprintf(w, "data: %s\n\n", msg)
|
||||
w.(http.Flusher).Flush()
|
||||
case <-ticker.C:
|
||||
fmt.Fprintf(w, ": keepalive\n\n")
|
||||
w.(http.Flusher).Flush()
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
func (s *Server) BroadcastSSE(event string) {
|
||||
s.sseClientsMu.Lock()
|
||||
defer s.sseClientsMu.Unlock()
|
||||
|
||||
for client := range s.sseClients {
|
||||
select {
|
||||
case client <- event:
|
||||
default:
|
||||
// Client channel full, skip
|
||||
}
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,503 @@
|
||||
// ===== CLIENT-SIDE ROUTER =====
|
||||
|
||||
let currentRoute = null;
|
||||
let pageIntervals = {}; // Track intervals per page for cleanup
|
||||
|
||||
const routes = {
|
||||
'/': 'overview',
|
||||
'/overview': 'overview',
|
||||
'/trades': 'trades',
|
||||
'/whitelist': 'whitelist',
|
||||
'/news': 'news'
|
||||
};
|
||||
|
||||
async function navigateTo(path, skipHistory = false) {
|
||||
const page = routes[path];
|
||||
if (!page) {
|
||||
console.warn('Unknown route:', path);
|
||||
return;
|
||||
}
|
||||
|
||||
// Skip if already on this page
|
||||
if (currentRoute === page) return;
|
||||
|
||||
// Update browser history
|
||||
if (!skipHistory) {
|
||||
history.pushState({ page }, '', path);
|
||||
}
|
||||
|
||||
// Clean up previous page
|
||||
cleanupPage(currentRoute);
|
||||
|
||||
// Update active tab
|
||||
updateActiveTab(page);
|
||||
|
||||
// Load new content
|
||||
await loadContent(page);
|
||||
|
||||
// Initialize new page
|
||||
initPage(page);
|
||||
|
||||
currentRoute = page;
|
||||
}
|
||||
|
||||
async function loadContent(page) {
|
||||
const contentDiv = document.getElementById('content');
|
||||
if (!contentDiv) return;
|
||||
|
||||
contentDiv.innerHTML = '<p style="text-align: center; padding: 40px;">Loading...</p>';
|
||||
|
||||
try {
|
||||
const resp = await fetch(`/content/${page}`);
|
||||
if (!resp.ok) throw new Error(`HTTP ${resp.status}`);
|
||||
|
||||
const html = await resp.text();
|
||||
contentDiv.innerHTML = html;
|
||||
} catch (err) {
|
||||
console.error('Failed to load content:', err);
|
||||
contentDiv.innerHTML = '<p style="text-align: center; padding: 40px; color: red;">Failed to load content</p>';
|
||||
}
|
||||
}
|
||||
|
||||
function updateActiveTab(page) {
|
||||
document.querySelectorAll('.tab').forEach(tab => {
|
||||
tab.classList.remove('active');
|
||||
});
|
||||
|
||||
const activeTab = document.querySelector(`.tab[data-page="${page}"]`);
|
||||
if (activeTab) {
|
||||
activeTab.classList.add('active');
|
||||
}
|
||||
}
|
||||
|
||||
function initPage(page) {
|
||||
switch(page) {
|
||||
case 'overview':
|
||||
loadBalance();
|
||||
loadTrades();
|
||||
break;
|
||||
case 'trades':
|
||||
loadTrades();
|
||||
break;
|
||||
case 'whitelist':
|
||||
loadWhitelist();
|
||||
break;
|
||||
case 'news':
|
||||
const newsInterval = initNewsPage();
|
||||
pageIntervals.newsRefresh = newsInterval;
|
||||
break;
|
||||
}
|
||||
}
|
||||
|
||||
function cleanupPage(page) {
|
||||
// Clear any intervals for previous page
|
||||
if (pageIntervals[page]) {
|
||||
clearInterval(pageIntervals[page]);
|
||||
delete pageIntervals[page];
|
||||
}
|
||||
}
|
||||
|
||||
// Intercept tab link clicks
|
||||
document.addEventListener('click', (e) => {
|
||||
const link = e.target.closest('a.tab');
|
||||
if (link && link.origin === location.origin) {
|
||||
e.preventDefault();
|
||||
const path = link.getAttribute('href');
|
||||
navigateTo(path);
|
||||
}
|
||||
});
|
||||
|
||||
// Handle browser back/forward
|
||||
window.addEventListener('popstate', (e) => {
|
||||
const path = location.pathname;
|
||||
navigateTo(path, true); // true = skip pushState
|
||||
});
|
||||
|
||||
// ===== NEWS PAGE FUNCTIONS =====
|
||||
|
||||
let newsCurrentFilter = 'all';
|
||||
let allNewsArticles = [];
|
||||
|
||||
function initNewsPage() {
|
||||
// Initialize news page
|
||||
loadNews();
|
||||
|
||||
// Set up filter buttons
|
||||
document.querySelectorAll('.filter-btn').forEach(btn => {
|
||||
btn.addEventListener('click', function() {
|
||||
document.querySelectorAll('.filter-btn').forEach(b => b.classList.remove('active'));
|
||||
this.classList.add('active');
|
||||
newsCurrentFilter = this.dataset.filter;
|
||||
renderNews(allNewsArticles);
|
||||
});
|
||||
});
|
||||
|
||||
// Auto-refresh every 30 seconds
|
||||
return setInterval(loadNews, 30000);
|
||||
}
|
||||
|
||||
function loadNews() {
|
||||
fetch('/api/news/recent?limit=50')
|
||||
.then(response => response.json())
|
||||
.then(data => {
|
||||
allNewsArticles = data || [];
|
||||
renderNews(allNewsArticles);
|
||||
})
|
||||
.catch(error => {
|
||||
console.error('Failed to load news:', error);
|
||||
const container = document.getElementById('news-list');
|
||||
if (container) {
|
||||
container.innerHTML = '<p class="no-news">Failed to load news. Please try again later.</p>';
|
||||
}
|
||||
});
|
||||
}
|
||||
|
||||
function renderNews(articles) {
|
||||
const container = document.getElementById('news-list');
|
||||
if (!container) return;
|
||||
|
||||
if (!articles || articles.length === 0) {
|
||||
container.innerHTML = '<p class="no-news">No news articles available</p>';
|
||||
return;
|
||||
}
|
||||
|
||||
const filtered = newsCurrentFilter === 'all'
|
||||
? articles
|
||||
: articles.filter(a => getSentimentLabel(a) === newsCurrentFilter);
|
||||
|
||||
if (filtered.length === 0) {
|
||||
container.innerHTML = '<p class="no-news">No ' + newsCurrentFilter + ' articles</p>';
|
||||
return;
|
||||
}
|
||||
|
||||
container.innerHTML = filtered.map(article => {
|
||||
const sentiment = getSentimentLabel(article);
|
||||
const score = article.SentimentScore;
|
||||
const symbols = article.Symbols ? article.Symbols.split(',').filter(s => s.trim()) : [];
|
||||
const method = article.SentimentMethod || 'keyword';
|
||||
|
||||
return `
|
||||
<div class="news-article ${sentiment}">
|
||||
<div class="news-article-header">
|
||||
<h3 class="news-title">
|
||||
<a href="${article.URL}" target="_blank" rel="noopener">${article.Title}</a>
|
||||
</h3>
|
||||
<div class="news-sentiment">
|
||||
${score !== null && score !== undefined ?
|
||||
`<span class="sentiment-score">${formatScore(score)}</span>` : ''}
|
||||
<span class="sentiment-badge ${sentiment}">${sentiment}</span>
|
||||
</div>
|
||||
</div>
|
||||
<div class="news-meta">
|
||||
<span class="news-source">${article.Source}</span>
|
||||
<span class="news-time" title="${new Date(article.PublishedAt).toLocaleString()}">${formatTimeAgo(article.PublishedAt)}</span>
|
||||
<span class="news-fetched" title="Fetched: ${new Date(article.FetchedAt).toLocaleString()}">📥 ${formatTimeAgo(article.FetchedAt)}</span>
|
||||
</div>
|
||||
${article.Content ? `<div class="news-content">${article.Content}</div>` : ''}
|
||||
${symbols.length > 0 ? `
|
||||
<div class="news-symbols">
|
||||
${symbols.map(s => `<span class="symbol-tag">${s.trim()}</span>`).join('')}
|
||||
</div>
|
||||
` : ''}
|
||||
<div class="news-method">Analysis: ${method}</div>
|
||||
</div>
|
||||
`;
|
||||
}).join('');
|
||||
|
||||
const countEl = document.getElementById('news-count');
|
||||
const updateEl = document.getElementById('last-update');
|
||||
if (countEl) countEl.textContent = filtered.length + ' article' + (filtered.length !== 1 ? 's' : '');
|
||||
if (updateEl) updateEl.textContent = 'Last update: ' + new Date().toLocaleTimeString();
|
||||
}
|
||||
|
||||
function formatTimeAgo(dateString) {
|
||||
const date = new Date(dateString);
|
||||
const now = new Date();
|
||||
const seconds = Math.floor((now - date) / 1000);
|
||||
|
||||
if (seconds < 60) return 'just now';
|
||||
if (seconds < 3600) return Math.floor(seconds / 60) + 'm ago';
|
||||
if (seconds < 86400) return Math.floor(seconds / 3600) + 'h ago';
|
||||
return Math.floor(seconds / 86400) + 'd ago';
|
||||
}
|
||||
|
||||
function formatScore(score) {
|
||||
if (score === null || score === undefined) return '';
|
||||
return (score > 0 ? '+' : '') + score.toFixed(2);
|
||||
}
|
||||
|
||||
function getSentimentLabel(article) {
|
||||
return article.SentimentLabel || 'neutral';
|
||||
}
|
||||
|
||||
// ===== SSE CONNECTION =====
|
||||
|
||||
// SSE connection
|
||||
const evtSource = new EventSource('/events');
|
||||
evtSource.onmessage = function(event) {
|
||||
console.log('SSE:', event.data);
|
||||
|
||||
// Update relevant data based on current page and event type
|
||||
if (event.data.includes('whitelist')) {
|
||||
if (currentRoute === 'whitelist') loadWhitelist();
|
||||
}
|
||||
|
||||
if (event.data.includes('trade')) {
|
||||
if (currentRoute === 'overview' || currentRoute === 'trades') {
|
||||
loadTrades();
|
||||
}
|
||||
}
|
||||
|
||||
if (event.data.includes('news')) {
|
||||
if (currentRoute === 'news') {
|
||||
console.log('News updated, reloading...');
|
||||
loadNews();
|
||||
}
|
||||
}
|
||||
|
||||
// Always update balance (shown in multiple places)
|
||||
loadBalance();
|
||||
};
|
||||
|
||||
// ===== INITIALIZATION =====
|
||||
|
||||
// Initial data and config (always load these)
|
||||
loadBalance();
|
||||
loadConfigOnce();
|
||||
|
||||
// Initial route dispatch
|
||||
if (window.initialPage) {
|
||||
// Server-rendered initial page, content already in DOM
|
||||
currentRoute = window.initialPage;
|
||||
updateActiveTab(currentRoute);
|
||||
initPage(currentRoute);
|
||||
} else {
|
||||
// Client-side navigation (e.g., page refresh)
|
||||
navigateTo(location.pathname, true);
|
||||
}
|
||||
|
||||
setInterval(() => { loadBalance(); loadTrades(); }, 10000);
|
||||
|
||||
async function loadConfigOnce() {
|
||||
// Load fresh config from API (no caching to ensure we always get latest state)
|
||||
try {
|
||||
const resp = await fetch('/api/config');
|
||||
const config = await resp.json();
|
||||
updateBanner(config);
|
||||
} catch (err) {
|
||||
console.error('Failed to load config:', err);
|
||||
}
|
||||
}
|
||||
|
||||
function updateBanner(config) {
|
||||
const banner = document.getElementById('mode-banner');
|
||||
if (banner) {
|
||||
if (config.dry_run) {
|
||||
banner.textContent = '⚠️ DRY RUN MODE';
|
||||
banner.className = 'mode-banner mode-dryrun loaded';
|
||||
} else {
|
||||
banner.textContent = '🔴 LIVE MODE';
|
||||
banner.className = 'mode-banner mode-live loaded';
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
async function loadConfig() {
|
||||
try {
|
||||
const resp = await fetch('/api/config');
|
||||
const config = await resp.json();
|
||||
|
||||
const banner = document.getElementById('mode-banner');
|
||||
if (banner) {
|
||||
if (config.dry_run) {
|
||||
banner.textContent = '⚠️ DRY RUN MODE';
|
||||
banner.className = 'mode-banner mode-dryrun loaded';
|
||||
} else {
|
||||
banner.textContent = '🔴 LIVE MODE';
|
||||
banner.className = 'mode-banner mode-live loaded';
|
||||
}
|
||||
}
|
||||
} catch (err) {
|
||||
console.error('Failed to load config:', err);
|
||||
}
|
||||
}
|
||||
|
||||
async function loadBalance() {
|
||||
const resp = await fetch('/api/balance');
|
||||
const data = await resp.json();
|
||||
const balanceEl = document.getElementById('balance');
|
||||
if (balanceEl) {
|
||||
balanceEl.textContent = '$' + (data.TotalValue || 0).toLocaleString('en-US', {
|
||||
minimumFractionDigits: 2,
|
||||
maximumFractionDigits: 2
|
||||
});
|
||||
}
|
||||
}
|
||||
|
||||
async function loadTrades() {
|
||||
const resp = await fetch('/api/trades');
|
||||
const trades = await resp.json();
|
||||
|
||||
let pending = 0, active = 0;
|
||||
|
||||
// Overview table (last 10)
|
||||
const overviewBody = document.getElementById('trades-body-overview');
|
||||
if (overviewBody) {
|
||||
overviewBody.innerHTML = '';
|
||||
trades.slice(0, 10).forEach(trade => {
|
||||
if (trade.status === 'PENDING') pending++;
|
||||
if (trade.status === 'SUBMITTED' || trade.status === 'FILLED') active++;
|
||||
overviewBody.innerHTML += buildTradeRow(trade, false);
|
||||
});
|
||||
}
|
||||
|
||||
// Full trades table
|
||||
const tradesBody = document.getElementById('trades-body');
|
||||
if (tradesBody) {
|
||||
tradesBody.innerHTML = '';
|
||||
trades.forEach(trade => {
|
||||
tradesBody.innerHTML += buildTradeRow(trade, true);
|
||||
});
|
||||
}
|
||||
|
||||
const pendingEl = document.getElementById('pending-trades');
|
||||
const activeEl = document.getElementById('active-trades');
|
||||
if (pendingEl) pendingEl.textContent = pending;
|
||||
if (activeEl) activeEl.textContent = active;
|
||||
}
|
||||
|
||||
function buildTradeRow(trade, detailed) {
|
||||
const dryBadge = trade.is_dry_run ? '<span class="dry-run-badge">DRY</span>' : '';
|
||||
const price = trade.executed_price ? '$' + trade.executed_price.toFixed(2) : '-';
|
||||
const pnl = formatPnL(trade.dry_run_pnl);
|
||||
const actions = trade.status === 'PENDING' ?
|
||||
'<button class="action-btn approve-btn" onclick="approveTrade(' + trade.id + ', false)">✓</button>' +
|
||||
'<button class="action-btn force-btn" onclick="approveTrade(' + trade.id + ', true)">⚡</button>' +
|
||||
'<button class="action-btn reject-btn" onclick="rejectTrade(' + trade.id + ')">✕</button>' : '-';
|
||||
|
||||
let row = '<tr>' +
|
||||
'<td>' + trade.id + dryBadge + '</td>' +
|
||||
'<td>' + trade.symbol + '</td>' +
|
||||
'<td>' + trade.action + '</td>' +
|
||||
'<td>' + trade.quantity + '</td>' +
|
||||
'<td><span class="status status-' + trade.status.toLowerCase() + '">' + trade.status + '</span></td>' +
|
||||
'<td>' + price + '</td>' +
|
||||
'<td>' + pnl + '</td>' +
|
||||
'<td>' + (trade.confidence * 100).toFixed(0) + '%</td>';
|
||||
|
||||
if (detailed) {
|
||||
row += '<td style="max-width: 200px; font-size: 11px;">' + (trade.reasoning || '-') + '</td>';
|
||||
}
|
||||
|
||||
row += '<td>' + new Date(trade.created_at).toLocaleString() + '</td>' +
|
||||
'<td>' + actions + '</td></tr>';
|
||||
|
||||
return row;
|
||||
}
|
||||
|
||||
function formatPnL(pnl) {
|
||||
if (!pnl) return '-';
|
||||
const formatted = '$' + Math.abs(pnl).toFixed(2);
|
||||
const cssClass = pnl >= 0 ? 'pnl-positive' : 'pnl-negative';
|
||||
const sign = pnl >= 0 ? '+' : '-';
|
||||
return '<span class="' + cssClass + '">' + sign + formatted + '</span>';
|
||||
}
|
||||
|
||||
async function approveTrade(id, forceNow) {
|
||||
await fetch('/api/trades/' + id + '/approve?force=' + forceNow, { method: 'POST' });
|
||||
loadTrades();
|
||||
}
|
||||
|
||||
async function rejectTrade(id) {
|
||||
const reason = prompt('Rejection reason:') || 'User rejected';
|
||||
await fetch('/api/trades/' + id + '/reject', {
|
||||
method: 'POST',
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify({ reason })
|
||||
});
|
||||
loadTrades();
|
||||
}
|
||||
|
||||
// Whitelist functions
|
||||
async function loadWhitelist() {
|
||||
const resp = await fetch('/api/whitelist');
|
||||
const entries = await resp.json();
|
||||
|
||||
const tbody = document.getElementById('whitelist-body');
|
||||
tbody.innerHTML = '';
|
||||
|
||||
entries.forEach(entry => {
|
||||
const statusClass = entry.Enabled ? 'enabled' : 'disabled';
|
||||
const statusText = entry.Enabled ? '✓ Enabled' : '✕ Disabled';
|
||||
tbody.innerHTML += '<tr>' +
|
||||
'<td><strong>' + entry.Symbol + '</strong></td>' +
|
||||
'<td>' + (entry.Name || '-') + '</td>' +
|
||||
'<td>' + (entry.WKN || '-') + '</td>' +
|
||||
'<td>' + (entry.ISIN || '-') + '</td>' +
|
||||
'<td class="' + statusClass + '">' + statusText + '</td>' +
|
||||
'<td style="max-width: 200px; font-size: 11px;">' + (entry.Notes || '-') + '</td>' +
|
||||
'<td>' +
|
||||
'<button class="action-btn edit-btn" onclick="editWhitelist(' + entry.ID + ')">✎</button>' +
|
||||
'<button class="action-btn delete-btn" onclick="deleteWhitelist(' + entry.ID + ')">🗑</button>' +
|
||||
'</td></tr>';
|
||||
});
|
||||
}
|
||||
|
||||
function showAddModal() {
|
||||
document.getElementById('modal-title').textContent = 'Add Symbol';
|
||||
document.getElementById('whitelist-form').reset();
|
||||
document.getElementById('entry-id').value = '';
|
||||
document.getElementById('entry-enabled').checked = true;
|
||||
document.getElementById('whitelist-modal').classList.add('active');
|
||||
}
|
||||
|
||||
async function editWhitelist(id) {
|
||||
const resp = await fetch('/api/whitelist/' + id);
|
||||
const entry = await resp.json();
|
||||
|
||||
document.getElementById('modal-title').textContent = 'Edit Symbol';
|
||||
document.getElementById('entry-id').value = entry.ID;
|
||||
document.getElementById('entry-symbol').value = entry.Symbol;
|
||||
document.getElementById('entry-name').value = entry.Name || '';
|
||||
document.getElementById('entry-wkn').value = entry.WKN || '';
|
||||
document.getElementById('entry-isin').value = entry.ISIN || '';
|
||||
document.getElementById('entry-enabled').checked = entry.Enabled;
|
||||
document.getElementById('entry-notes').value = entry.Notes || '';
|
||||
document.getElementById('whitelist-modal').classList.add('active');
|
||||
}
|
||||
|
||||
async function saveWhitelist(event) {
|
||||
event.preventDefault();
|
||||
|
||||
const id = document.getElementById('entry-id').value;
|
||||
const data = {
|
||||
symbol: document.getElementById('entry-symbol').value,
|
||||
name: document.getElementById('entry-name').value,
|
||||
wkn: document.getElementById('entry-wkn').value,
|
||||
isin: document.getElementById('entry-isin').value,
|
||||
enabled: document.getElementById('entry-enabled').checked,
|
||||
notes: document.getElementById('entry-notes').value
|
||||
};
|
||||
|
||||
const url = id ? '/api/whitelist/' + id : '/api/whitelist';
|
||||
const method = id ? 'PUT' : 'POST';
|
||||
|
||||
await fetch(url, {
|
||||
method: method,
|
||||
headers: { 'Content-Type': 'application/json' },
|
||||
body: JSON.stringify(data)
|
||||
});
|
||||
|
||||
closeModal();
|
||||
loadWhitelist();
|
||||
}
|
||||
|
||||
async function deleteWhitelist(id) {
|
||||
if (!confirm('Delete this symbol from whitelist?')) return;
|
||||
await fetch('/api/whitelist/' + id, { method: 'DELETE' });
|
||||
loadWhitelist();
|
||||
}
|
||||
|
||||
function closeModal() {
|
||||
document.getElementById('whitelist-modal').classList.remove('active');
|
||||
}
|
||||
@@ -0,0 +1,442 @@
|
||||
body {
|
||||
font-family: Arial, sans-serif;
|
||||
margin: 20px;
|
||||
background: #f5f5f5;
|
||||
}
|
||||
|
||||
.container {
|
||||
max-width: 1600px;
|
||||
margin: 0 auto;
|
||||
background: white;
|
||||
padding: 20px;
|
||||
border-radius: 8px;
|
||||
}
|
||||
|
||||
h1 {
|
||||
color: #333;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
|
||||
h2 {
|
||||
color: #555;
|
||||
margin-top: 30px;
|
||||
border-bottom: 2px solid #ddd;
|
||||
padding-bottom: 5px;
|
||||
}
|
||||
|
||||
.tabs {
|
||||
display: flex;
|
||||
border-bottom: 2px solid #ddd;
|
||||
margin-bottom: 20px;
|
||||
}
|
||||
|
||||
.tab {
|
||||
padding: 10px 20px;
|
||||
cursor: pointer;
|
||||
background: #f0f0f0;
|
||||
border: none;
|
||||
margin-right: 5px;
|
||||
text-decoration: none;
|
||||
color: inherit;
|
||||
display: inline-block;
|
||||
}
|
||||
|
||||
.tab:hover {
|
||||
background: #e0e0e0;
|
||||
}
|
||||
|
||||
.tab.active {
|
||||
background: white;
|
||||
border-bottom: 3px solid #4CAF50;
|
||||
}
|
||||
|
||||
.tab-content {
|
||||
display: block;
|
||||
}
|
||||
|
||||
.mode-banner {
|
||||
padding: 6px 16px;
|
||||
margin: 0;
|
||||
border-radius: 20px;
|
||||
font-weight: bold;
|
||||
text-align: center;
|
||||
font-size: 0.85em;
|
||||
opacity: 0;
|
||||
transition: opacity 0.3s ease;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.mode-banner.loaded {
|
||||
opacity: 1;
|
||||
}
|
||||
|
||||
.mode-live {
|
||||
background: #f8d7da;
|
||||
color: #721c24;
|
||||
border: 2px solid #dc3545;
|
||||
}
|
||||
|
||||
.mode-dryrun {
|
||||
background: #fff3cd;
|
||||
color: #856404;
|
||||
border: 2px solid #ffc107;
|
||||
}
|
||||
|
||||
.stats {
|
||||
display: grid;
|
||||
grid-template-columns: repeat(3, 1fr);
|
||||
gap: 20px;
|
||||
margin: 20px 0;
|
||||
}
|
||||
|
||||
.stat-card {
|
||||
background: #f9f9f9;
|
||||
padding: 15px;
|
||||
border-radius: 5px;
|
||||
border-left: 4px solid #4CAF50;
|
||||
}
|
||||
|
||||
.stat-card h3 {
|
||||
margin: 0 0 10px 0;
|
||||
color: #666;
|
||||
font-size: 14px;
|
||||
}
|
||||
|
||||
.stat-card .value {
|
||||
font-size: 24px;
|
||||
font-weight: bold;
|
||||
color: #333;
|
||||
}
|
||||
|
||||
table {
|
||||
width: 100%;
|
||||
border-collapse: collapse;
|
||||
margin-top: 10px;
|
||||
}
|
||||
|
||||
th, td {
|
||||
padding: 12px;
|
||||
text-align: left;
|
||||
border-bottom: 1px solid #ddd;
|
||||
}
|
||||
|
||||
th {
|
||||
background: #f0f0f0;
|
||||
font-weight: bold;
|
||||
}
|
||||
|
||||
.status {
|
||||
padding: 4px 8px;
|
||||
border-radius: 3px;
|
||||
font-size: 12px;
|
||||
font-weight: bold;
|
||||
}
|
||||
|
||||
.status-pending { background: #FFF3CD; color: #856404; }
|
||||
.status-approved { background: #D1ECF1; color: #0C5460; }
|
||||
.status-submitted { background: #D4EDDA; color: #155724; }
|
||||
.status-completed { background: #D4EDDA; color: #155724; }
|
||||
.status-rejected { background: #F8D7DA; color: #721C24; }
|
||||
|
||||
.action-btn {
|
||||
padding: 6px 12px;
|
||||
margin: 2px;
|
||||
border: none;
|
||||
border-radius: 3px;
|
||||
cursor: pointer;
|
||||
font-size: 12px;
|
||||
}
|
||||
|
||||
.approve-btn { background: #28a745; color: white; }
|
||||
.reject-btn { background: #dc3545; color: white; }
|
||||
.force-btn { background: #ffc107; color: black; }
|
||||
.add-btn { background: #007bff; color: white; padding: 10px 20px; border: none; border-radius: 5px; cursor: pointer; }
|
||||
.edit-btn { background: #17a2b8; color: white; }
|
||||
.delete-btn { background: #dc3545; color: white; }
|
||||
.save-btn { background: #28a745; color: white; padding: 8px 16px; }
|
||||
.cancel-btn { background: #6c757d; color: white; padding: 8px 16px; }
|
||||
|
||||
.enabled { color: #28a745; font-weight: bold; }
|
||||
.disabled { color: #dc3545; font-weight: bold; }
|
||||
.dry-run-badge { background: #ffc107; color: black; padding: 2px 6px; border-radius: 3px; font-size: 11px; margin-left: 5px; }
|
||||
.pnl-positive { color: #28a745; font-weight: bold; }
|
||||
.pnl-negative { color: #dc3545; font-weight: bold; }
|
||||
|
||||
.modal {
|
||||
display: none;
|
||||
position: fixed;
|
||||
top: 0;
|
||||
left: 0;
|
||||
width: 100%;
|
||||
height: 100%;
|
||||
background: rgba(0,0,0,0.5);
|
||||
z-index: 1000;
|
||||
}
|
||||
|
||||
.modal.active {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
justify-content: center;
|
||||
}
|
||||
|
||||
.modal-content {
|
||||
background: white;
|
||||
padding: 30px;
|
||||
border-radius: 8px;
|
||||
max-width: 500px;
|
||||
width: 90%;
|
||||
}
|
||||
|
||||
.modal-content h2 {
|
||||
margin-top: 0;
|
||||
}
|
||||
|
||||
.form-group {
|
||||
margin-bottom: 15px;
|
||||
}
|
||||
|
||||
.form-group label {
|
||||
display: block;
|
||||
margin-bottom: 5px;
|
||||
font-weight: bold;
|
||||
}
|
||||
|
||||
.form-group input, .form-group textarea {
|
||||
width: 100%;
|
||||
padding: 8px;
|
||||
border: 1px solid #ddd;
|
||||
border-radius: 4px;
|
||||
}
|
||||
|
||||
.form-group textarea {
|
||||
height: 80px;
|
||||
resize: vertical;
|
||||
}
|
||||
|
||||
.form-actions {
|
||||
display: flex;
|
||||
gap: 10px;
|
||||
margin-top: 20px;
|
||||
}
|
||||
|
||||
.checkbox-label {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
}
|
||||
|
||||
.checkbox-label input[type="checkbox"] {
|
||||
width: auto;
|
||||
}
|
||||
|
||||
/* ===== NEWS PAGE STYLES ===== */
|
||||
|
||||
.news-container {
|
||||
max-width: 1200px;
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.news-header {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
margin-bottom: 20px;
|
||||
padding-bottom: 10px;
|
||||
border-bottom: 2px solid #333;
|
||||
}
|
||||
|
||||
.news-stats {
|
||||
display: flex;
|
||||
gap: 20px;
|
||||
font-size: 0.9em;
|
||||
color: #666;
|
||||
}
|
||||
|
||||
.news-filters {
|
||||
display: flex;
|
||||
gap: 10px;
|
||||
margin-bottom: 20px;
|
||||
}
|
||||
|
||||
.filter-btn {
|
||||
padding: 8px 16px;
|
||||
border: 2px solid #ddd;
|
||||
background: white;
|
||||
cursor: pointer;
|
||||
border-radius: 4px;
|
||||
transition: all 0.2s;
|
||||
}
|
||||
|
||||
.filter-btn:hover {
|
||||
background: #f0f0f0;
|
||||
}
|
||||
|
||||
.filter-btn.active {
|
||||
background: #007bff;
|
||||
color: white;
|
||||
border-color: #007bff;
|
||||
}
|
||||
|
||||
.news-list {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 15px;
|
||||
}
|
||||
|
||||
.news-article {
|
||||
padding: 15px;
|
||||
border-radius: 8px;
|
||||
border-left: 4px solid #ccc;
|
||||
background: white;
|
||||
box-shadow: 0 2px 4px rgba(0,0,0,0.1);
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
|
||||
.news-article:hover {
|
||||
transform: translateX(5px);
|
||||
}
|
||||
|
||||
.news-article.positive {
|
||||
border-left-color: #28a745;
|
||||
background: #f0fff4;
|
||||
}
|
||||
|
||||
.news-article.negative {
|
||||
border-left-color: #dc3545;
|
||||
background: #fff5f5;
|
||||
}
|
||||
|
||||
.news-article.neutral {
|
||||
border-left-color: #6c757d;
|
||||
background: #f8f9fa;
|
||||
}
|
||||
|
||||
.news-article-header {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: start;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
|
||||
.news-title {
|
||||
font-size: 1.1em;
|
||||
font-weight: bold;
|
||||
margin: 0;
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
.news-title a {
|
||||
color: #333;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
.news-title a:hover {
|
||||
color: #007bff;
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
.news-sentiment {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
margin-left: 15px;
|
||||
}
|
||||
|
||||
.sentiment-badge {
|
||||
padding: 4px 12px;
|
||||
border-radius: 12px;
|
||||
font-size: 0.85em;
|
||||
font-weight: bold;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.sentiment-badge.positive {
|
||||
background: #28a745;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-badge.negative {
|
||||
background: #dc3545;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-badge.neutral {
|
||||
background: #6c757d;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-score {
|
||||
font-size: 0.9em;
|
||||
color: #666;
|
||||
font-weight: bold;
|
||||
}
|
||||
|
||||
.news-meta {
|
||||
display: flex;
|
||||
gap: 15px;
|
||||
font-size: 0.85em;
|
||||
color: #666;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
|
||||
.news-source {
|
||||
font-weight: bold;
|
||||
color: #007bff;
|
||||
}
|
||||
|
||||
.news-time {
|
||||
color: #999;
|
||||
}
|
||||
|
||||
.news-fetched {
|
||||
color: #999;
|
||||
font-size: 0.9em;
|
||||
}
|
||||
|
||||
.news-fetched:hover {
|
||||
color: #666;
|
||||
cursor: help;
|
||||
}
|
||||
|
||||
.news-content {
|
||||
color: #555;
|
||||
line-height: 1.5;
|
||||
margin-top: 8px;
|
||||
}
|
||||
|
||||
.news-symbols {
|
||||
margin-top: 10px;
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.symbol-tag {
|
||||
padding: 3px 8px;
|
||||
background: #e9ecef;
|
||||
border-radius: 4px;
|
||||
font-size: 0.85em;
|
||||
font-weight: bold;
|
||||
color: #495057;
|
||||
}
|
||||
|
||||
.news-method {
|
||||
font-size: 0.8em;
|
||||
color: #999;
|
||||
font-style: italic;
|
||||
margin-top: 5px;
|
||||
}
|
||||
|
||||
.loading {
|
||||
text-align: center;
|
||||
padding: 40px;
|
||||
color: #999;
|
||||
}
|
||||
|
||||
.no-news {
|
||||
text-align: center;
|
||||
padding: 40px;
|
||||
color: #999;
|
||||
font-style: italic;
|
||||
}
|
||||
|
||||
@@ -0,0 +1,14 @@
|
||||
<!DOCTYPE html>
|
||||
<html>
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<title>{{.Title}}</title>
|
||||
<link rel="stylesheet" href="/static/style.css">
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
{{template "content" .}}
|
||||
</div>
|
||||
<script src="/static/app.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,24 @@
|
||||
{{define "content"}}
|
||||
<div style="display: flex; align-items: center; justify-content: space-between; margin-bottom: 10px;">
|
||||
<h1 style="margin: 0;">🤖 AI Trading Dashboard</h1>
|
||||
<div id="mode-banner" class="mode-banner mode-dryrun"></div>
|
||||
</div>
|
||||
|
||||
<div class="tabs">
|
||||
<a href="/overview" class="tab {{if eq .CurrentPage "overview"}}active{{end}}">Overview</a>
|
||||
<a href="/trades" class="tab {{if eq .CurrentPage "trades"}}active{{end}}">Trades</a>
|
||||
<a href="/whitelist" class="tab {{if eq .CurrentPage "whitelist"}}active{{end}}">Whitelist</a>
|
||||
<a href="/news" class="tab {{if eq .CurrentPage "news"}}active{{end}}">News</a>
|
||||
</div>
|
||||
|
||||
{{if eq .CurrentPage "overview"}}
|
||||
{{template "overview" .}}
|
||||
{{else if eq .CurrentPage "trades"}}
|
||||
{{template "trades" .}}
|
||||
{{else if eq .CurrentPage "whitelist"}}
|
||||
{{template "whitelist" .}}
|
||||
{{template "whitelist-modal" .}}
|
||||
{{else}}
|
||||
{{template "overview" .}}
|
||||
{{end}}
|
||||
{{end}}
|
||||
@@ -0,0 +1,22 @@
|
||||
{{define "news-content"}}
|
||||
<div class="news-container">
|
||||
<div class="news-header">
|
||||
<h1>📰 Live News Feed</h1>
|
||||
<div class="news-stats">
|
||||
<span id="news-count">0 articles</span>
|
||||
<span id="last-update">Last update: never</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="news-filters">
|
||||
<button class="filter-btn active" data-filter="all">All</button>
|
||||
<button class="filter-btn" data-filter="positive">Positive</button>
|
||||
<button class="filter-btn" data-filter="neutral">Neutral</button>
|
||||
<button class="filter-btn" data-filter="negative">Negative</button>
|
||||
</div>
|
||||
|
||||
<div id="news-list" class="news-list">
|
||||
<p class="loading">Loading news...</p>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,344 @@
|
||||
<!DOCTYPE html>
|
||||
<html>
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<title>{{.Title}}</title>
|
||||
<link rel="stylesheet" href="/static/style.css">
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
<div class="news-container">
|
||||
<div class="news-header">
|
||||
<h1>📰 Live News Feed</h1>
|
||||
<div class="news-stats">
|
||||
<span id="news-count">0 articles</span>
|
||||
<span id="last-update">Last update: never</span>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<div class="news-filters">
|
||||
<button class="filter-btn active" data-filter="all">All</button>
|
||||
<button class="filter-btn" data-filter="positive">Positive</button>
|
||||
<button class="filter-btn" data-filter="neutral">Neutral</button>
|
||||
<button class="filter-btn" data-filter="negative">Negative</button>
|
||||
</div>
|
||||
|
||||
<div id="news-list" class="news-list">
|
||||
<p class="loading">Loading news...</p>
|
||||
</div>
|
||||
</div>
|
||||
</div>
|
||||
<script src="/static/app.js"></script>
|
||||
|
||||
<style>
|
||||
.news-container {
|
||||
max-width: 1200px;
|
||||
margin: 0 auto;
|
||||
}
|
||||
|
||||
.news-header {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: center;
|
||||
margin-bottom: 20px;
|
||||
padding-bottom: 10px;
|
||||
border-bottom: 2px solid #333;
|
||||
}
|
||||
|
||||
.news-stats {
|
||||
display: flex;
|
||||
gap: 20px;
|
||||
font-size: 0.9em;
|
||||
color: #666;
|
||||
}
|
||||
|
||||
.news-filters {
|
||||
display: flex;
|
||||
gap: 10px;
|
||||
margin-bottom: 20px;
|
||||
}
|
||||
|
||||
.filter-btn {
|
||||
padding: 8px 16px;
|
||||
border: 2px solid #ddd;
|
||||
background: white;
|
||||
cursor: pointer;
|
||||
border-radius: 4px;
|
||||
transition: all 0.2s;
|
||||
}
|
||||
|
||||
.filter-btn:hover {
|
||||
background: #f0f0f0;
|
||||
}
|
||||
|
||||
.filter-btn.active {
|
||||
background: #007bff;
|
||||
color: white;
|
||||
border-color: #007bff;
|
||||
}
|
||||
|
||||
.news-list {
|
||||
display: flex;
|
||||
flex-direction: column;
|
||||
gap: 15px;
|
||||
}
|
||||
|
||||
.news-article {
|
||||
padding: 15px;
|
||||
border-radius: 8px;
|
||||
border-left: 4px solid #ccc;
|
||||
background: white;
|
||||
box-shadow: 0 2px 4px rgba(0,0,0,0.1);
|
||||
transition: transform 0.2s;
|
||||
}
|
||||
|
||||
.news-article:hover {
|
||||
transform: translateX(5px);
|
||||
}
|
||||
|
||||
.news-article.positive {
|
||||
border-left-color: #28a745;
|
||||
background: #f0fff4;
|
||||
}
|
||||
|
||||
.news-article.negative {
|
||||
border-left-color: #dc3545;
|
||||
background: #fff5f5;
|
||||
}
|
||||
|
||||
.news-article.neutral {
|
||||
border-left-color: #6c757d;
|
||||
background: #f8f9fa;
|
||||
}
|
||||
|
||||
.news-article-header {
|
||||
display: flex;
|
||||
justify-content: space-between;
|
||||
align-items: start;
|
||||
margin-bottom: 10px;
|
||||
}
|
||||
|
||||
.news-title {
|
||||
font-size: 1.1em;
|
||||
font-weight: bold;
|
||||
margin: 0;
|
||||
flex: 1;
|
||||
}
|
||||
|
||||
.news-title a {
|
||||
color: #333;
|
||||
text-decoration: none;
|
||||
}
|
||||
|
||||
.news-title a:hover {
|
||||
color: #007bff;
|
||||
text-decoration: underline;
|
||||
}
|
||||
|
||||
.news-sentiment {
|
||||
display: flex;
|
||||
align-items: center;
|
||||
gap: 8px;
|
||||
margin-left: 15px;
|
||||
}
|
||||
|
||||
.sentiment-badge {
|
||||
padding: 4px 12px;
|
||||
border-radius: 12px;
|
||||
font-size: 0.85em;
|
||||
font-weight: bold;
|
||||
white-space: nowrap;
|
||||
}
|
||||
|
||||
.sentiment-badge.positive {
|
||||
background: #28a745;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-badge.negative {
|
||||
background: #dc3545;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-badge.neutral {
|
||||
background: #6c757d;
|
||||
color: white;
|
||||
}
|
||||
|
||||
.sentiment-score {
|
||||
font-size: 0.9em;
|
||||
color: #666;
|
||||
font-weight: bold;
|
||||
}
|
||||
|
||||
.news-meta {
|
||||
display: flex;
|
||||
gap: 15px;
|
||||
font-size: 0.85em;
|
||||
color: #666;
|
||||
margin-bottom: 8px;
|
||||
}
|
||||
|
||||
.news-source {
|
||||
font-weight: bold;
|
||||
color: #007bff;
|
||||
}
|
||||
|
||||
.news-time {
|
||||
color: #999;
|
||||
}
|
||||
|
||||
.news-content {
|
||||
color: #555;
|
||||
line-height: 1.5;
|
||||
margin-top: 8px;
|
||||
}
|
||||
|
||||
.news-symbols {
|
||||
margin-top: 10px;
|
||||
display: flex;
|
||||
gap: 8px;
|
||||
flex-wrap: wrap;
|
||||
}
|
||||
|
||||
.symbol-tag {
|
||||
padding: 3px 8px;
|
||||
background: #e9ecef;
|
||||
border-radius: 4px;
|
||||
font-size: 0.85em;
|
||||
font-weight: bold;
|
||||
color: #495057;
|
||||
}
|
||||
|
||||
.news-method {
|
||||
font-size: 0.8em;
|
||||
color: #999;
|
||||
font-style: italic;
|
||||
margin-top: 5px;
|
||||
}
|
||||
|
||||
.loading {
|
||||
text-align: center;
|
||||
padding: 40px;
|
||||
color: #999;
|
||||
}
|
||||
|
||||
.no-news {
|
||||
text-align: center;
|
||||
padding: 40px;
|
||||
color: #999;
|
||||
font-style: italic;
|
||||
}
|
||||
</style>
|
||||
|
||||
<script>
|
||||
let currentFilter = 'all';
|
||||
let allNews = [];
|
||||
|
||||
function formatTimeAgo(dateString) {
|
||||
const date = new Date(dateString);
|
||||
const now = new Date();
|
||||
const seconds = Math.floor((now - date) / 1000);
|
||||
|
||||
if (seconds < 60) return 'just now';
|
||||
if (seconds < 3600) return Math.floor(seconds / 60) + 'm ago';
|
||||
if (seconds < 86400) return Math.floor(seconds / 3600) + 'h ago';
|
||||
return Math.floor(seconds / 86400) + 'd ago';
|
||||
}
|
||||
|
||||
function formatScore(score) {
|
||||
if (score === null || score === undefined) return '';
|
||||
return (score > 0 ? '+' : '') + score.toFixed(2);
|
||||
}
|
||||
|
||||
function getSentimentLabel(article) {
|
||||
return article.SentimentLabel || 'neutral';
|
||||
}
|
||||
|
||||
function renderNews(articles) {
|
||||
const container = document.getElementById('news-list');
|
||||
|
||||
if (!articles || articles.length === 0) {
|
||||
container.innerHTML = '<p class="no-news">No news articles available</p>';
|
||||
return;
|
||||
}
|
||||
|
||||
const filtered = currentFilter === 'all'
|
||||
? articles
|
||||
: articles.filter(a => getSentimentLabel(a) === currentFilter);
|
||||
|
||||
if (filtered.length === 0) {
|
||||
container.innerHTML = '<p class="no-news">No ' + currentFilter + ' articles</p>';
|
||||
return;
|
||||
}
|
||||
|
||||
container.innerHTML = filtered.map(article => {
|
||||
const sentiment = getSentimentLabel(article);
|
||||
const score = article.SentimentScore;
|
||||
const symbols = article.Symbols ? article.Symbols.split(',').filter(s => s.trim()) : [];
|
||||
const method = article.SentimentMethod || 'keyword';
|
||||
|
||||
return `
|
||||
<div class="news-article ${sentiment}">
|
||||
<div class="news-article-header">
|
||||
<h3 class="news-title">
|
||||
<a href="${article.URL}" target="_blank" rel="noopener">${article.Title}</a>
|
||||
</h3>
|
||||
<div class="news-sentiment">
|
||||
${score !== null && score !== undefined ?
|
||||
`<span class="sentiment-score">${formatScore(score)}</span>` : ''}
|
||||
<span class="sentiment-badge ${sentiment}">${sentiment}</span>
|
||||
</div>
|
||||
</div>
|
||||
<div class="news-meta">
|
||||
<span class="news-source">${article.Source}</span>
|
||||
<span class="news-time">${formatTimeAgo(article.PublishedAt)}</span>
|
||||
</div>
|
||||
${article.Content ? `<div class="news-content">${article.Content}</div>` : ''}
|
||||
${symbols.length > 0 ? `
|
||||
<div class="news-symbols">
|
||||
${symbols.map(s => `<span class="symbol-tag">${s.trim()}</span>`).join('')}
|
||||
</div>
|
||||
` : ''}
|
||||
<div class="news-method">Analysis: ${method}</div>
|
||||
</div>
|
||||
`;
|
||||
}).join('');
|
||||
|
||||
// Update stats
|
||||
document.getElementById('news-count').textContent = filtered.length + ' article' + (filtered.length !== 1 ? 's' : '');
|
||||
document.getElementById('last-update').textContent = 'Last update: ' + new Date().toLocaleTimeString();
|
||||
}
|
||||
|
||||
function loadNews() {
|
||||
fetch('/api/news/recent?limit=50')
|
||||
.then(response => response.json())
|
||||
.then(data => {
|
||||
allNews = data || [];
|
||||
renderNews(allNews);
|
||||
})
|
||||
.catch(error => {
|
||||
console.error('Failed to load news:', error);
|
||||
document.getElementById('news-list').innerHTML =
|
||||
'<p class="no-news">Failed to load news. Please try again later.</p>';
|
||||
});
|
||||
}
|
||||
|
||||
// Filter buttons
|
||||
document.querySelectorAll('.filter-btn').forEach(btn => {
|
||||
btn.addEventListener('click', function() {
|
||||
document.querySelectorAll('.filter-btn').forEach(b => b.classList.remove('active'));
|
||||
this.classList.add('active');
|
||||
currentFilter = this.dataset.filter;
|
||||
renderNews(allNews);
|
||||
});
|
||||
});
|
||||
|
||||
// Initial load
|
||||
loadNews();
|
||||
|
||||
// Auto-refresh every 30 seconds
|
||||
setInterval(loadNews, 30000);
|
||||
</script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,39 @@
|
||||
{{define "overview-content"}}
|
||||
<div id="tab-overview" class="tab-content">
|
||||
<div class="stats">
|
||||
<div class="stat-card">
|
||||
<h3>Total Balance</h3>
|
||||
<div class="value" id="balance">$0.00</div>
|
||||
</div>
|
||||
<div class="stat-card">
|
||||
<h3>Active Trades</h3>
|
||||
<div class="value" id="active-trades">0</div>
|
||||
</div>
|
||||
<div class="stat-card">
|
||||
<h3>Pending Trades</h3>
|
||||
<div class="value" id="pending-trades">0</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<h2>Recent Trades</h2>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>ID</th>
|
||||
<th>Symbol</th>
|
||||
<th>Action</th>
|
||||
<th>Qty</th>
|
||||
<th>Status</th>
|
||||
<th>Price</th>
|
||||
<th>P&L</th>
|
||||
<th>Confidence</th>
|
||||
<th>Created</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="trades-body-overview">
|
||||
<tr><td colspan="10" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,39 @@
|
||||
{{define "overview"}}
|
||||
<div id="tab-overview" class="tab-content">
|
||||
<div class="stats">
|
||||
<div class="stat-card">
|
||||
<h3>Total Balance</h3>
|
||||
<div class="value" id="balance">$0.00</div>
|
||||
</div>
|
||||
<div class="stat-card">
|
||||
<h3>Active Trades</h3>
|
||||
<div class="value" id="active-trades">0</div>
|
||||
</div>
|
||||
<div class="stat-card">
|
||||
<h3>Pending Trades</h3>
|
||||
<div class="value" id="pending-trades">0</div>
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<h2>Recent Trades</h2>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>ID</th>
|
||||
<th>Symbol</th>
|
||||
<th>Action</th>
|
||||
<th>Qty</th>
|
||||
<th>Status</th>
|
||||
<th>Price</th>
|
||||
<th>P&L</th>
|
||||
<th>Confidence</th>
|
||||
<th>Created</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="trades-body-overview">
|
||||
<tr><td colspan="10" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,40 @@
|
||||
<!DOCTYPE html>
|
||||
<html>
|
||||
<head>
|
||||
<meta charset="UTF-8">
|
||||
<title>{{.Title}}</title>
|
||||
<link rel="stylesheet" href="/static/style.css">
|
||||
</head>
|
||||
<body>
|
||||
<div class="container">
|
||||
<div style="display: flex; align-items: center; justify-content: space-between; margin-bottom: 10px;">
|
||||
<h1 style="margin: 0;">🤖 AI Trading Dashboard</h1>
|
||||
<div id="mode-banner" class="mode-banner mode-dryrun"></div>
|
||||
</div>
|
||||
|
||||
<div class="tabs">
|
||||
<a href="/overview" class="tab{{if eq .CurrentPage "overview"}} active{{end}}" data-page="overview">Overview</a>
|
||||
<a href="/trades" class="tab{{if eq .CurrentPage "trades"}} active{{end}}" data-page="trades">Trades</a>
|
||||
<a href="/whitelist" class="tab{{if eq .CurrentPage "whitelist"}} active{{end}}" data-page="whitelist">Whitelist</a>
|
||||
<a href="/news" class="tab{{if eq .CurrentPage "news"}} active{{end}}" data-page="news">News</a>
|
||||
</div>
|
||||
|
||||
<div id="content">
|
||||
{{if eq .CurrentPage "overview"}}
|
||||
{{template "overview-content" .}}
|
||||
{{else if eq .CurrentPage "trades"}}
|
||||
{{template "trades-content" .}}
|
||||
{{else if eq .CurrentPage "whitelist"}}
|
||||
{{template "whitelist-content" .}}
|
||||
{{else if eq .CurrentPage "news"}}
|
||||
{{template "news-content" .}}
|
||||
{{else}}
|
||||
{{template "overview-content" .}}
|
||||
{{end}}
|
||||
</div>
|
||||
</div>
|
||||
|
||||
<script>window.initialPage = "{{.CurrentPage}}";</script>
|
||||
<script src="/static/app.js"></script>
|
||||
</body>
|
||||
</html>
|
||||
@@ -0,0 +1,25 @@
|
||||
{{define "trades-content"}}
|
||||
<div id="tab-trades" class="tab-content">
|
||||
<h2>All Trades</h2>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>ID</th>
|
||||
<th>Symbol</th>
|
||||
<th>Action</th>
|
||||
<th>Quantity</th>
|
||||
<th>Status</th>
|
||||
<th>Price</th>
|
||||
<th>P&L</th>
|
||||
<th>Confidence</th>
|
||||
<th>Reasoning</th>
|
||||
<th>Created</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="trades-body">
|
||||
<tr><td colspan="11" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,25 @@
|
||||
{{define "trades"}}
|
||||
<div id="tab-trades" class="tab-content">
|
||||
<h2>All Trades</h2>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>ID</th>
|
||||
<th>Symbol</th>
|
||||
<th>Action</th>
|
||||
<th>Quantity</th>
|
||||
<th>Status</th>
|
||||
<th>Price</th>
|
||||
<th>P&L</th>
|
||||
<th>Confidence</th>
|
||||
<th>Reasoning</th>
|
||||
<th>Created</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="trades-body">
|
||||
<tr><td colspan="11" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,63 @@
|
||||
{{define "whitelist-content"}}
|
||||
<div id="tab-whitelist" class="tab-content">
|
||||
<div style="display: flex; justify-content: space-between; align-items: center;">
|
||||
<h2 style="margin: 0;">Trading Whitelist</h2>
|
||||
<button class="add-btn" onclick="showAddModal()">➕ Add Symbol</button>
|
||||
</div>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Symbol</th>
|
||||
<th>Name</th>
|
||||
<th>WKN</th>
|
||||
<th>ISIN</th>
|
||||
<th>Status</th>
|
||||
<th>Notes</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="whitelist-body">
|
||||
<tr><td colspan="7" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
|
||||
<div id="whitelist-modal" class="modal">
|
||||
<div class="modal-content">
|
||||
<h2 id="modal-title">Add Symbol</h2>
|
||||
<form id="whitelist-form" onsubmit="saveWhitelist(event)">
|
||||
<input type="hidden" id="entry-id" value="">
|
||||
<div class="form-group">
|
||||
<label>Symbol *</label>
|
||||
<input type="text" id="entry-symbol" required>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>Name</label>
|
||||
<input type="text" id="entry-name">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>WKN</label>
|
||||
<input type="text" id="entry-wkn">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>ISIN</label>
|
||||
<input type="text" id="entry-isin">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label class="checkbox-label">
|
||||
<input type="checkbox" id="entry-enabled" checked>
|
||||
<span>Enabled (allow trading)</span>
|
||||
</label>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>Notes</label>
|
||||
<textarea id="entry-notes"></textarea>
|
||||
</div>
|
||||
<div class="form-actions">
|
||||
<button type="submit" class="action-btn save-btn">Save</button>
|
||||
<button type="button" class="action-btn cancel-btn" onclick="closeModal()">Cancel</button>
|
||||
</div>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,40 @@
|
||||
{{define "whitelist-modal"}}
|
||||
<div id="whitelist-modal" class="modal">
|
||||
<div class="modal-content">
|
||||
<h2 id="modal-title">Add Symbol</h2>
|
||||
<form id="whitelist-form" onsubmit="saveWhitelist(event)">
|
||||
<input type="hidden" id="entry-id" value="">
|
||||
<div class="form-group">
|
||||
<label>Symbol *</label>
|
||||
<input type="text" id="entry-symbol" required>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>Name</label>
|
||||
<input type="text" id="entry-name">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>WKN</label>
|
||||
<input type="text" id="entry-wkn">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>ISIN</label>
|
||||
<input type="text" id="entry-isin">
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label class="checkbox-label">
|
||||
<input type="checkbox" id="entry-enabled" checked>
|
||||
<span>Enabled (allow trading)</span>
|
||||
</label>
|
||||
</div>
|
||||
<div class="form-group">
|
||||
<label>Notes</label>
|
||||
<textarea id="entry-notes"></textarea>
|
||||
</div>
|
||||
<div class="form-actions">
|
||||
<button type="submit" class="action-btn save-btn">Save</button>
|
||||
<button type="button" class="action-btn cancel-btn" onclick="closeModal()">Cancel</button>
|
||||
</div>
|
||||
</form>
|
||||
</div>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,24 @@
|
||||
{{define "whitelist"}}
|
||||
<div id="tab-whitelist" class="tab-content">
|
||||
<div style="display: flex; justify-content: space-between; align-items: center;">
|
||||
<h2 style="margin: 0;">Trading Whitelist</h2>
|
||||
<button class="add-btn" onclick="showAddModal()">➕ Add Symbol</button>
|
||||
</div>
|
||||
<table>
|
||||
<thead>
|
||||
<tr>
|
||||
<th>Symbol</th>
|
||||
<th>Name</th>
|
||||
<th>WKN</th>
|
||||
<th>ISIN</th>
|
||||
<th>Status</th>
|
||||
<th>Notes</th>
|
||||
<th>Actions</th>
|
||||
</tr>
|
||||
</thead>
|
||||
<tbody id="whitelist-body">
|
||||
<tr><td colspan="7" style="text-align: center;">Loading...</td></tr>
|
||||
</tbody>
|
||||
</table>
|
||||
</div>
|
||||
{{end}}
|
||||
@@ -0,0 +1,238 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
type Config struct {
|
||||
IBGateway IBGatewayConfig `yaml:"ib_gateway"`
|
||||
Trading TradingConfig `yaml:"trading"`
|
||||
Database DatabaseConfig `yaml:"database"`
|
||||
Web WebConfig `yaml:"web"`
|
||||
OIDC OIDCConfig `yaml:"oidc"`
|
||||
News NewsConfig `yaml:"news"`
|
||||
LLMScorer LLMScorerConfig `yaml:"llm_scorer"`
|
||||
LogLevel string `yaml:"log_level"`
|
||||
}
|
||||
|
||||
type IBGatewayConfig struct {
|
||||
Host string `yaml:"host"`
|
||||
Port int `yaml:"port"`
|
||||
ClientID int `yaml:"client_id"`
|
||||
MarketDataType int `yaml:"market_data_type"` // 1=Live, 2=Frozen, 3=Delayed(15min), 4=Delayed-Frozen
|
||||
}
|
||||
|
||||
type TradingConfig struct {
|
||||
Strategy string `yaml:"strategy"`
|
||||
StopLossEnabled bool `yaml:"stop_loss_enabled"`
|
||||
StopLossPercent float64 `yaml:"stop_loss_percent"`
|
||||
MaxTradesPerHour int `yaml:"max_trades_per_hour"`
|
||||
MaxParallelTrades int `yaml:"max_parallel_trades"`
|
||||
PendingTime Duration `yaml:"pending_time"`
|
||||
DryRun bool `yaml:"dry_run"`
|
||||
DryRunBalance float64 `yaml:"dry_run_balance"`
|
||||
TradingEnabled bool `yaml:"trading_enabled"`
|
||||
TradingInterval Duration `yaml:"trading_interval"`
|
||||
WatchSymbols []string `yaml:"watch_symbols"`
|
||||
TakeProfitPercent float64 `yaml:"take_profit_percent"`
|
||||
HoldTimeMinutes int `yaml:"hold_time_minutes"`
|
||||
NegativeSentiment bool `yaml:"sell_on_negative_sentiment"`
|
||||
MaxTradeValue float64 `yaml:"max_trade_value"`
|
||||
}
|
||||
|
||||
type DatabaseConfig struct {
|
||||
Path string `yaml:"path"`
|
||||
}
|
||||
|
||||
type WebConfig struct {
|
||||
Port string `yaml:"port"`
|
||||
Host string `yaml:"host"`
|
||||
}
|
||||
|
||||
type OIDCConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
Issuer string `yaml:"issuer"`
|
||||
ClientID string `yaml:"client_id"`
|
||||
ClientSecret string `yaml:"client_secret"`
|
||||
RedirectURL string `yaml:"redirect_url"`
|
||||
Scopes []string `yaml:"scopes"`
|
||||
}
|
||||
|
||||
type NewsConfig struct {
|
||||
PollInterval Duration `yaml:"poll_interval"`
|
||||
DefaultRateLimit *RateLimit `yaml:"default_rate_limit,omitempty"`
|
||||
Sources []NewsSource `yaml:"sources"`
|
||||
}
|
||||
|
||||
type NewsSource struct {
|
||||
Name string `yaml:"name"`
|
||||
URL string `yaml:"url"`
|
||||
Type string `yaml:"type"` // "rss", "api", etc.
|
||||
Enabled bool `yaml:"enabled"`
|
||||
Auth *NewsAuth `yaml:"auth,omitempty"`
|
||||
Headers map[string]string `yaml:"headers,omitempty"`
|
||||
RateLimit *RateLimit `yaml:"rate_limit,omitempty"`
|
||||
}
|
||||
|
||||
type RateLimit struct {
|
||||
MaxPerHour int `yaml:"max_per_hour,omitempty"` // Max requests per hour, 0 = unlimited
|
||||
MaxPerDay int `yaml:"max_per_day,omitempty"` // Max requests per day, 0 = unlimited
|
||||
}
|
||||
|
||||
type NewsAuth struct {
|
||||
Type string `yaml:"type"` // "basic", "bearer", "apikey"
|
||||
Username string `yaml:"username,omitempty"`
|
||||
Password string `yaml:"password,omitempty"`
|
||||
Token string `yaml:"token,omitempty"`
|
||||
}
|
||||
|
||||
type LLMScorerConfig struct {
|
||||
Enabled bool `yaml:"enabled"`
|
||||
Endpoint string `yaml:"endpoint"`
|
||||
ModelName string `yaml:"model_name"`
|
||||
Timeout Duration `yaml:"timeout"`
|
||||
Temperature float64 `yaml:"temperature"`
|
||||
MaxRetries int `yaml:"max_retries"`
|
||||
EnsembleWeight float64 `yaml:"ensemble_weight"` // 0.0-1.0, 1.0 = LLM only
|
||||
}
|
||||
|
||||
func Load() (*Config, error) {
|
||||
cfg := &Config{
|
||||
IBGateway: IBGatewayConfig{
|
||||
Host: getEnv("IB_GATEWAY_HOST", "127.0.0.1"),
|
||||
Port: getEnvInt("IB_GATEWAY_PORT", 4001),
|
||||
ClientID: getEnvInt("IB_CLIENT_ID", 1),
|
||||
MarketDataType: getEnvInt("IB_MARKET_DATA_TYPE", 3), // Default: 3 = Delayed (15min, kostenlos)
|
||||
},
|
||||
Trading: TradingConfig{
|
||||
Strategy: getEnv("TRADING_STRATEGY", "normal"),
|
||||
StopLossEnabled: getEnvBool("STOP_LOSS_ENABLED", true),
|
||||
StopLossPercent: getEnvFloat("STOP_LOSS_PERCENT", 3.0),
|
||||
MaxTradesPerHour: getEnvInt("MAX_TRADES_PER_HOUR", 6),
|
||||
MaxParallelTrades: getEnvInt("MAX_PARALLEL_TRADES", 5),
|
||||
PendingTime: Duration{
|
||||
Duration: time.Duration(getEnvInt("PENDING_TIME_SECONDS", 300)) * time.Second,
|
||||
},
|
||||
DryRun: getEnvBool("DRY_RUN", true),
|
||||
DryRunBalance: getEnvFloat("DRY_RUN_BALANCE", 100000.0),
|
||||
TradingEnabled: getEnvBool("TRADING_ENABLED", true),
|
||||
TradingInterval: Duration{
|
||||
Duration: time.Duration(getEnvInt("TRADING_INTERVAL_SECONDS", 60)) * time.Second,
|
||||
},
|
||||
WatchSymbols: strings.Split(getEnv("WATCH_SYMBOLS", "AAPL,MSFT,GOOGL,TSLA,AMZN"), ","),
|
||||
TakeProfitPercent: getEnvFloat("TAKE_PROFIT_PERCENT", 5.0),
|
||||
HoldTimeMinutes: getEnvInt("HOLD_TIME_MINUTES", 30),
|
||||
NegativeSentiment: getEnvBool("SELL_ON_NEGATIVE_SENTIMENT", true),
|
||||
MaxTradeValue: getEnvFloat("MAX_TRADE_VALUE", 0.0),
|
||||
},
|
||||
Database: DatabaseConfig{
|
||||
Path: getEnv("DB_PATH", "./data/aitrade.db"),
|
||||
},
|
||||
Web: WebConfig{
|
||||
Port: getEnv("WEB_PORT", "8080"),
|
||||
Host: getEnv("WEB_HOST", "0.0.0.0"),
|
||||
},
|
||||
OIDC: OIDCConfig{
|
||||
Enabled: getEnvBool("OIDC_ENABLED", false),
|
||||
Issuer: getEnv("OIDC_ISSUER", ""),
|
||||
ClientID: getEnv("OIDC_CLIENT_ID", ""),
|
||||
ClientSecret: getEnv("OIDC_CLIENT_SECRET", ""),
|
||||
RedirectURL: getEnv("OIDC_REDIRECT_URL", ""),
|
||||
Scopes: strings.Split(getEnv("OIDC_SCOPES", "openid,profile,email"), ","),
|
||||
},
|
||||
News: NewsConfig{
|
||||
PollInterval: Duration{
|
||||
Duration: time.Duration(getEnvInt("NEWS_POLL_INTERVAL", 300)) * time.Second,
|
||||
},
|
||||
Sources: []NewsSource{}, // Will be loaded from YAML
|
||||
},
|
||||
LLMScorer: LLMScorerConfig{
|
||||
Enabled: getEnvBool("LLM_SCORER_ENABLED", false),
|
||||
Endpoint: getEnv("LLM_SCORER_ENDPOINT", "http://localhost:11434"),
|
||||
ModelName: getEnv("LLM_SCORER_MODEL", "mistral"),
|
||||
Timeout: Duration{
|
||||
Duration: time.Duration(getEnvInt("LLM_SCORER_TIMEOUT_SECONDS", 30)) * time.Second,
|
||||
},
|
||||
Temperature: getEnvFloat("LLM_SCORER_TEMPERATURE", 0.3),
|
||||
MaxRetries: getEnvInt("LLM_SCORER_MAX_RETRIES", 2),
|
||||
EnsembleWeight: getEnvFloat("LLM_SCORER_ENSEMBLE_WEIGHT", 0.7),
|
||||
},
|
||||
LogLevel: getEnv("LOG_LEVEL", "info"),
|
||||
}
|
||||
|
||||
if err := cfg.Validate(); err != nil {
|
||||
return nil, err
|
||||
}
|
||||
|
||||
return cfg, nil
|
||||
}
|
||||
|
||||
func (c *Config) Validate() error {
|
||||
validStrategies := map[string]bool{
|
||||
"defensive": true,
|
||||
"normal": true,
|
||||
"aggressive": true,
|
||||
}
|
||||
|
||||
if !validStrategies[c.Trading.Strategy] {
|
||||
return fmt.Errorf("invalid strategy: %s (must be defensive, normal, or aggressive)", c.Trading.Strategy)
|
||||
}
|
||||
|
||||
if c.Trading.StopLossPercent < 0 || c.Trading.StopLossPercent > 100 {
|
||||
return fmt.Errorf("stop loss percent must be between 0 and 100")
|
||||
}
|
||||
|
||||
if c.Trading.MaxTradesPerHour < 1 {
|
||||
return fmt.Errorf("max trades per hour must be at least 1")
|
||||
}
|
||||
|
||||
if c.Trading.MaxParallelTrades < 1 {
|
||||
return fmt.Errorf("max parallel trades must be at least 1")
|
||||
}
|
||||
|
||||
if c.OIDC.Enabled {
|
||||
if c.OIDC.Issuer == "" || c.OIDC.ClientID == "" || c.OIDC.ClientSecret == "" {
|
||||
return fmt.Errorf("OIDC enabled but missing required config (issuer, client_id, or client_secret)")
|
||||
}
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func getEnv(key, defaultValue string) string {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
return value
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
func getEnvInt(key string, defaultValue int) int {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
if intVal, err := strconv.Atoi(value); err == nil {
|
||||
return intVal
|
||||
}
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
func getEnvFloat(key string, defaultValue float64) float64 {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
if floatVal, err := strconv.ParseFloat(value, 64); err == nil {
|
||||
return floatVal
|
||||
}
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
|
||||
func getEnvBool(key string, defaultValue bool) bool {
|
||||
if value := os.Getenv(key); value != "" {
|
||||
if boolVal, err := strconv.ParseBool(value); err == nil {
|
||||
return boolVal
|
||||
}
|
||||
}
|
||||
return defaultValue
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"strconv"
|
||||
"strings"
|
||||
"time"
|
||||
)
|
||||
|
||||
// Duration wraps time.Duration to support YAML unmarshaling from seconds
|
||||
type Duration struct {
|
||||
time.Duration
|
||||
}
|
||||
|
||||
// UnmarshalYAML implements yaml.Unmarshaler interface
|
||||
// Accepts either:
|
||||
// - Integer: interpreted as seconds (e.g., 300)
|
||||
// - String: Go duration format (e.g., "5m", "1h30m")
|
||||
func (d *Duration) UnmarshalYAML(unmarshal func(interface{}) error) error {
|
||||
// Try unmarshaling as integer (seconds)
|
||||
var seconds int
|
||||
if err := unmarshal(&seconds); err == nil {
|
||||
d.Duration = time.Duration(seconds) * time.Second
|
||||
return nil
|
||||
}
|
||||
|
||||
// Try unmarshaling as string (Go duration format)
|
||||
var str string
|
||||
if err := unmarshal(&str); err != nil {
|
||||
return err
|
||||
}
|
||||
|
||||
// Parse as Go duration ("5m", "1h30m", etc.)
|
||||
parsed, err := time.ParseDuration(str)
|
||||
if err != nil {
|
||||
// If not a valid Go duration, try parsing as number + "s" suffix
|
||||
if strings.HasSuffix(str, "s") {
|
||||
numStr := strings.TrimSuffix(str, "s")
|
||||
if sec, err := strconv.Atoi(numStr); err == nil {
|
||||
d.Duration = time.Duration(sec) * time.Second
|
||||
return nil
|
||||
}
|
||||
}
|
||||
return fmt.Errorf("invalid duration format: %s (use seconds or Go duration like '5m')", str)
|
||||
}
|
||||
|
||||
d.Duration = parsed
|
||||
return nil
|
||||
}
|
||||
|
||||
// MarshalYAML implements yaml.Marshaler interface
|
||||
func (d Duration) MarshalYAML() (interface{}, error) {
|
||||
// Always marshal as seconds for consistency
|
||||
return int(d.Duration.Seconds()), nil
|
||||
}
|
||||
@@ -0,0 +1,55 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"log/slog"
|
||||
"os"
|
||||
|
||||
"github.com/rs/zerolog"
|
||||
slogzerolog "github.com/samber/slog-zerolog"
|
||||
"github.com/scmhub/ibapi"
|
||||
)
|
||||
|
||||
// SetupLogger creates a logger based on the config and sets up ibapi logging
|
||||
func (c *Config) SetupLogger() *slog.Logger {
|
||||
// Parse log level
|
||||
logLevel := parseLogLevel(c.LogLevel)
|
||||
|
||||
// Create shared zerolog logger with config level
|
||||
zerolog.SetGlobalLevel(slogToZerologLevel(logLevel))
|
||||
zerologLogger := zerolog.New(os.Stdout).With().Timestamp().Logger()
|
||||
|
||||
// Set it for ibapi
|
||||
ibapi.SetLogger(zerologLogger)
|
||||
|
||||
// Use it as backend for slog
|
||||
return slog.New(slogzerolog.Option{
|
||||
Level: logLevel,
|
||||
Logger: &zerologLogger,
|
||||
}.NewZerologHandler())
|
||||
}
|
||||
|
||||
func parseLogLevel(level string) slog.Level {
|
||||
switch level {
|
||||
case "debug":
|
||||
return slog.LevelDebug
|
||||
case "warn":
|
||||
return slog.LevelWarn
|
||||
case "error":
|
||||
return slog.LevelError
|
||||
default:
|
||||
return slog.LevelInfo
|
||||
}
|
||||
}
|
||||
|
||||
func slogToZerologLevel(level slog.Level) zerolog.Level {
|
||||
switch level {
|
||||
case slog.LevelDebug:
|
||||
return zerolog.DebugLevel
|
||||
case slog.LevelWarn:
|
||||
return zerolog.WarnLevel
|
||||
case slog.LevelError:
|
||||
return zerolog.ErrorLevel
|
||||
default:
|
||||
return zerolog.InfoLevel
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,147 @@
|
||||
package config
|
||||
|
||||
import (
|
||||
"fmt"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"time"
|
||||
|
||||
"gopkg.in/yaml.v3"
|
||||
)
|
||||
|
||||
// LoadYAMLOrEnv loads configuration from YAML file or falls back to environment variables
|
||||
func LoadYAMLOrEnv() (*Config, error) {
|
||||
// Check for config file (in order of preference)
|
||||
configPaths := []string{
|
||||
os.Getenv("CONFIG_FILE"), // Highest priority
|
||||
"./config.yaml",
|
||||
"./config.yml",
|
||||
filepath.Join(os.Getenv("HOME"), ".config/aitrade/config.yaml"),
|
||||
"/etc/aitrade/config.yaml",
|
||||
}
|
||||
|
||||
for _, path := range configPaths {
|
||||
if path == "" {
|
||||
continue
|
||||
}
|
||||
|
||||
if _, err := os.Stat(path); err == nil {
|
||||
cfg, err := loadYAMLFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to load %s: %w", path, err)
|
||||
}
|
||||
fmt.Printf("Loaded configuration from: %s\n", path)
|
||||
return cfg, nil
|
||||
}
|
||||
}
|
||||
|
||||
// Fall back to environment variables
|
||||
return Load()
|
||||
}
|
||||
|
||||
// loadYAMLFile loads and parses a YAML configuration file
|
||||
func loadYAMLFile(path string) (*Config, error) {
|
||||
data, err := os.ReadFile(path)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to read file: %w", err)
|
||||
}
|
||||
|
||||
var cfg Config
|
||||
if err := yaml.Unmarshal(data, &cfg); err != nil {
|
||||
return nil, fmt.Errorf("failed to parse YAML: %w", err)
|
||||
}
|
||||
|
||||
// Apply defaults for empty values
|
||||
applyDefaults(&cfg)
|
||||
|
||||
return &cfg, cfg.Validate()
|
||||
}
|
||||
|
||||
// applyDefaults sets default values for empty fields
|
||||
func applyDefaults(cfg *Config) {
|
||||
if cfg.IBGateway.Host == "" {
|
||||
cfg.IBGateway.Host = "127.0.0.1"
|
||||
}
|
||||
if cfg.IBGateway.Port == 0 {
|
||||
cfg.IBGateway.Port = 4001
|
||||
}
|
||||
if cfg.IBGateway.ClientID == 0 {
|
||||
cfg.IBGateway.ClientID = 1
|
||||
}
|
||||
if cfg.IBGateway.MarketDataType == 0 {
|
||||
cfg.IBGateway.MarketDataType = 3 // Default: Delayed (15min, kostenlos)
|
||||
}
|
||||
|
||||
if cfg.Trading.Strategy == "" {
|
||||
cfg.Trading.Strategy = "normal"
|
||||
}
|
||||
if cfg.Trading.StopLossPercent == 0.0 {
|
||||
cfg.Trading.StopLossPercent = 3.0
|
||||
}
|
||||
if cfg.Trading.MaxTradesPerHour == 0 {
|
||||
cfg.Trading.MaxTradesPerHour = 6
|
||||
}
|
||||
if cfg.Trading.MaxParallelTrades == 0 {
|
||||
cfg.Trading.MaxParallelTrades = 5
|
||||
}
|
||||
if cfg.Trading.PendingTime.Duration == 0 {
|
||||
cfg.Trading.PendingTime.Duration = 300 * time.Second
|
||||
}
|
||||
if cfg.Trading.DryRunBalance == 0.0 {
|
||||
cfg.Trading.DryRunBalance = 100000.0
|
||||
}
|
||||
if cfg.Trading.TradingInterval.Duration == 0 {
|
||||
cfg.Trading.TradingInterval.Duration = 60 * time.Second
|
||||
}
|
||||
if len(cfg.Trading.WatchSymbols) == 0 {
|
||||
cfg.Trading.WatchSymbols = []string{"AAPL", "MSFT", "GOOGL", "TSLA", "AMZN"}
|
||||
}
|
||||
if cfg.Trading.TakeProfitPercent == 0.0 {
|
||||
cfg.Trading.TakeProfitPercent = 5.0
|
||||
}
|
||||
if cfg.Trading.HoldTimeMinutes == 0 {
|
||||
cfg.Trading.HoldTimeMinutes = 30
|
||||
}
|
||||
|
||||
if cfg.Database.Path == "" {
|
||||
cfg.Database.Path = "./data/aitrade.db"
|
||||
}
|
||||
|
||||
if cfg.Web.Port == "" {
|
||||
cfg.Web.Port = "8080"
|
||||
}
|
||||
if cfg.Web.Host == "" {
|
||||
cfg.Web.Host = "0.0.0.0"
|
||||
}
|
||||
|
||||
if len(cfg.OIDC.Scopes) == 0 {
|
||||
cfg.OIDC.Scopes = []string{"openid", "profile", "email"}
|
||||
}
|
||||
|
||||
if cfg.News.PollInterval.Duration == 0 {
|
||||
cfg.News.PollInterval.Duration = 300 * time.Second
|
||||
}
|
||||
|
||||
if cfg.LLMScorer.Endpoint == "" {
|
||||
cfg.LLMScorer.Endpoint = "http://localhost:11434"
|
||||
}
|
||||
if cfg.LLMScorer.ModelName == "" {
|
||||
cfg.LLMScorer.ModelName = "mistral"
|
||||
}
|
||||
if cfg.LLMScorer.Timeout.Duration == 0 {
|
||||
cfg.LLMScorer.Timeout.Duration = 30 * time.Second
|
||||
}
|
||||
if cfg.LLMScorer.Temperature == 0.0 {
|
||||
cfg.LLMScorer.Temperature = 0.3
|
||||
}
|
||||
if cfg.LLMScorer.MaxRetries == 0 {
|
||||
cfg.LLMScorer.MaxRetries = 2
|
||||
}
|
||||
if cfg.LLMScorer.EnsembleWeight == 0.0 {
|
||||
cfg.LLMScorer.EnsembleWeight = 0.7
|
||||
}
|
||||
|
||||
if cfg.LogLevel == "" {
|
||||
cfg.LogLevel = "info"
|
||||
}
|
||||
}
|
||||
@@ -0,0 +1,113 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type BalanceRepository struct {
|
||||
db Database
|
||||
}
|
||||
|
||||
func NewBalanceRepository(db Database) *BalanceRepository {
|
||||
return &BalanceRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *BalanceRepository) Create(ctx context.Context, balance *model.Balance) error {
|
||||
query := `
|
||||
INSERT INTO balances (timestamp, total_value, cash_balance, buying_power, unrealized_pnl, realized_pnl)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
`
|
||||
|
||||
result, err := r.db.ExecContext(ctx, query,
|
||||
balance.Timestamp,
|
||||
balance.TotalValue,
|
||||
balance.CashBalance,
|
||||
balance.BuyingPower,
|
||||
balance.UnrealizedPnL,
|
||||
balance.RealizedPnL,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert balance: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get last insert id: %w", err)
|
||||
}
|
||||
|
||||
balance.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *BalanceRepository) GetLatest(ctx context.Context) (*model.Balance, error) {
|
||||
query := `
|
||||
SELECT id, timestamp, total_value, cash_balance, buying_power, unrealized_pnl, realized_pnl
|
||||
FROM balances
|
||||
ORDER BY timestamp DESC
|
||||
LIMIT 1
|
||||
`
|
||||
|
||||
var balance model.Balance
|
||||
err := r.db.QueryRowContext(ctx, query).Scan(
|
||||
&balance.ID,
|
||||
&balance.Timestamp,
|
||||
&balance.TotalValue,
|
||||
&balance.CashBalance,
|
||||
&balance.BuyingPower,
|
||||
&balance.UnrealizedPnL,
|
||||
&balance.RealizedPnL,
|
||||
)
|
||||
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil
|
||||
}
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get latest balance: %w", err)
|
||||
}
|
||||
|
||||
return &balance, nil
|
||||
}
|
||||
|
||||
func (r *BalanceRepository) GetHistory(ctx context.Context, since time.Time, limit int) ([]*model.Balance, error) {
|
||||
query := `
|
||||
SELECT id, timestamp, total_value, cash_balance, buying_power, unrealized_pnl, realized_pnl
|
||||
FROM balances
|
||||
WHERE timestamp >= ?
|
||||
ORDER BY timestamp DESC
|
||||
LIMIT ?
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query, since, limit)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query balance history: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var balances []*model.Balance
|
||||
for rows.Next() {
|
||||
var balance model.Balance
|
||||
if err := rows.Scan(
|
||||
&balance.ID,
|
||||
&balance.Timestamp,
|
||||
&balance.TotalValue,
|
||||
&balance.CashBalance,
|
||||
&balance.BuyingPower,
|
||||
&balance.UnrealizedPnL,
|
||||
&balance.RealizedPnL,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan balance: %w", err)
|
||||
}
|
||||
balances = append(balances, &balance)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating balance rows: %w", err)
|
||||
}
|
||||
|
||||
return balances, nil
|
||||
}
|
||||
+185
@@ -0,0 +1,185 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"embed"
|
||||
"fmt"
|
||||
"io/fs"
|
||||
"log/slog"
|
||||
"os"
|
||||
"path/filepath"
|
||||
"sort"
|
||||
|
||||
_ "modernc.org/sqlite"
|
||||
)
|
||||
|
||||
//go:embed migrations/*.sql
|
||||
var migrationsFS embed.FS
|
||||
|
||||
// Database is the interface for database operations
|
||||
type Database interface {
|
||||
ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error)
|
||||
QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error)
|
||||
QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row
|
||||
BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error)
|
||||
Close() error
|
||||
}
|
||||
|
||||
type database struct {
|
||||
db *sql.DB
|
||||
logger *slog.Logger
|
||||
}
|
||||
|
||||
func New(dbPath string, logger *slog.Logger) (Database, error) {
|
||||
// Ensure data directory exists
|
||||
if err := os.MkdirAll(filepath.Dir(dbPath), 0755); err != nil {
|
||||
return nil, fmt.Errorf("failed to create data directory: %w", err)
|
||||
}
|
||||
|
||||
// Open database
|
||||
// Note: Using modernc.org/sqlite (pure Go, no CGO) instead of mattn/go-sqlite3
|
||||
db, err := sql.Open("sqlite", dbPath+"?_pragma=busy_timeout(5000)")
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to open database: %w", err)
|
||||
}
|
||||
|
||||
// Test connection
|
||||
if err := db.Ping(); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("failed to ping database: %w", err)
|
||||
}
|
||||
|
||||
dbInstance := &database{
|
||||
db: db,
|
||||
logger: logger,
|
||||
}
|
||||
|
||||
// Run migrations
|
||||
if err := dbInstance.migrate(); err != nil {
|
||||
db.Close()
|
||||
return nil, fmt.Errorf("failed to migrate: %w", err)
|
||||
}
|
||||
|
||||
return dbInstance, nil
|
||||
}
|
||||
|
||||
func (d *database) migrate() error {
|
||||
d.logger.Info("running database migrations")
|
||||
|
||||
// Create schema_migrations table if not exists
|
||||
_, err := d.db.Exec(`
|
||||
CREATE TABLE IF NOT EXISTS schema_migrations (
|
||||
version INTEGER PRIMARY KEY,
|
||||
applied_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP
|
||||
)
|
||||
`)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to create schema_migrations table: %w", err)
|
||||
}
|
||||
|
||||
// Get already applied migrations
|
||||
appliedMigrations := make(map[int]bool)
|
||||
rows, err := d.db.Query("SELECT version FROM schema_migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to query schema_migrations: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
for rows.Next() {
|
||||
var version int
|
||||
if err := rows.Scan(&version); err != nil {
|
||||
return fmt.Errorf("failed to scan version: %w", err)
|
||||
}
|
||||
appliedMigrations[version] = true
|
||||
}
|
||||
|
||||
// Read migrations from embedded filesystem
|
||||
entries, err := fs.ReadDir(migrationsFS, "migrations")
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read embedded migrations: %w", err)
|
||||
}
|
||||
|
||||
// Sort migration files by name
|
||||
var files []string
|
||||
for _, entry := range entries {
|
||||
if !entry.IsDir() && filepath.Ext(entry.Name()) == ".sql" {
|
||||
files = append(files, entry.Name())
|
||||
}
|
||||
}
|
||||
sort.Strings(files)
|
||||
|
||||
// Execute each migration
|
||||
for _, filename := range files {
|
||||
// Extract version from filename (e.g., "001_init.sql" -> 1)
|
||||
var version int
|
||||
if _, err := fmt.Sscanf(filename, "%d_", &version); err != nil {
|
||||
d.logger.Warn("skipping migration with invalid name format", slog.String("file", filename))
|
||||
continue
|
||||
}
|
||||
|
||||
// Skip if already applied
|
||||
if appliedMigrations[version] {
|
||||
d.logger.Debug("skipping already applied migration", slog.String("file", filename), slog.Int("version", version))
|
||||
continue
|
||||
}
|
||||
|
||||
// Read migration file from embedded FS
|
||||
content, err := migrationsFS.ReadFile(filepath.Join("migrations", filename))
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to read migration %s: %w", filename, err)
|
||||
}
|
||||
|
||||
// Execute migration in a transaction
|
||||
tx, err := d.db.Begin()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to begin transaction for migration %s: %w", filename, err)
|
||||
}
|
||||
|
||||
if _, err := tx.Exec(string(content)); err != nil {
|
||||
tx.Rollback()
|
||||
return fmt.Errorf("failed to execute migration %s: %w", filename, err)
|
||||
}
|
||||
|
||||
// Record migration
|
||||
if _, err := tx.Exec("INSERT INTO schema_migrations (version) VALUES (?)", version); err != nil {
|
||||
tx.Rollback()
|
||||
return fmt.Errorf("failed to record migration %s: %w", filename, err)
|
||||
}
|
||||
|
||||
if err := tx.Commit(); err != nil {
|
||||
return fmt.Errorf("failed to commit migration %s: %w", filename, err)
|
||||
}
|
||||
|
||||
d.logger.Info("applied migration", slog.String("file", filename), slog.Int("version", version))
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (d *database) Close() error {
|
||||
if d.db != nil {
|
||||
return d.db.Close()
|
||||
}
|
||||
return nil
|
||||
}
|
||||
|
||||
// ExecContext implements Database interface
|
||||
func (d *database) ExecContext(ctx context.Context, query string, args ...any) (sql.Result, error) {
|
||||
return d.db.ExecContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
// QueryContext implements Database interface
|
||||
func (d *database) QueryContext(ctx context.Context, query string, args ...any) (*sql.Rows, error) {
|
||||
return d.db.QueryContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
// QueryRowContext implements Database interface
|
||||
func (d *database) QueryRowContext(ctx context.Context, query string, args ...any) *sql.Row {
|
||||
return d.db.QueryRowContext(ctx, query, args...)
|
||||
}
|
||||
|
||||
// BeginTx implements Database interface
|
||||
func (d *database) BeginTx(ctx context.Context, opts *sql.TxOptions) (*sql.Tx, error) {
|
||||
return d.db.BeginTx(ctx, opts)
|
||||
}
|
||||
@@ -0,0 +1,61 @@
|
||||
-- Initial schema for AI Trading Application
|
||||
|
||||
-- Trades table
|
||||
CREATE TABLE IF NOT EXISTS trades (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
symbol TEXT NOT NULL,
|
||||
action TEXT NOT NULL CHECK(action IN ('BUY', 'SELL')),
|
||||
quantity INTEGER NOT NULL,
|
||||
status TEXT NOT NULL CHECK(status IN ('PENDING', 'APPROVED', 'REJECTED', 'SUBMITTED', 'FILLED', 'COMPLETED', 'STOPPED')),
|
||||
confidence REAL NOT NULL,
|
||||
reasoning TEXT,
|
||||
|
||||
target_price REAL,
|
||||
executed_price REAL,
|
||||
stop_loss_price REAL,
|
||||
ib_order_id INTEGER,
|
||||
|
||||
created_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
pending_until TIMESTAMP,
|
||||
approved_at TIMESTAMP,
|
||||
rejected_at TIMESTAMP,
|
||||
submitted_at TIMESTAMP,
|
||||
filled_at TIMESTAMP,
|
||||
completed_at TIMESTAMP,
|
||||
|
||||
rejection_reason TEXT,
|
||||
forced_by_user BOOLEAN DEFAULT FALSE
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_trades_status ON trades(status);
|
||||
CREATE INDEX IF NOT EXISTS idx_trades_symbol ON trades(symbol);
|
||||
CREATE INDEX IF NOT EXISTS idx_trades_created_at ON trades(created_at DESC);
|
||||
|
||||
-- Balances table
|
||||
CREATE TABLE IF NOT EXISTS balances (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
timestamp TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
total_value REAL NOT NULL,
|
||||
cash_balance REAL NOT NULL,
|
||||
buying_power REAL NOT NULL,
|
||||
unrealized_pnl REAL,
|
||||
realized_pnl REAL,
|
||||
|
||||
UNIQUE(timestamp)
|
||||
);
|
||||
|
||||
-- News articles table
|
||||
CREATE TABLE IF NOT EXISTS news_articles (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
source TEXT NOT NULL,
|
||||
title TEXT NOT NULL,
|
||||
url TEXT NOT NULL UNIQUE,
|
||||
content TEXT,
|
||||
published_at TIMESTAMP NOT NULL,
|
||||
fetched_at TIMESTAMP NOT NULL DEFAULT CURRENT_TIMESTAMP,
|
||||
symbols TEXT,
|
||||
sentiment_score REAL,
|
||||
sentiment_label TEXT
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_news_published_at ON news_articles(published_at DESC);
|
||||
@@ -0,0 +1,7 @@
|
||||
-- Add dry run support
|
||||
|
||||
ALTER TABLE trades ADD COLUMN is_dry_run BOOLEAN DEFAULT FALSE;
|
||||
ALTER TABLE trades ADD COLUMN dry_run_pnl REAL;
|
||||
|
||||
-- Add index for dry run trades
|
||||
CREATE INDEX IF NOT EXISTS idx_trades_dry_run ON trades(is_dry_run);
|
||||
@@ -0,0 +1,23 @@
|
||||
-- Whitelist table for approved trading symbols
|
||||
CREATE TABLE IF NOT EXISTS whitelist (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
symbol TEXT NOT NULL UNIQUE,
|
||||
name TEXT,
|
||||
wkn TEXT,
|
||||
isin TEXT,
|
||||
enabled BOOLEAN NOT NULL DEFAULT TRUE,
|
||||
created_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
notes TEXT
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_whitelist_symbol ON whitelist(symbol);
|
||||
CREATE INDEX IF NOT EXISTS idx_whitelist_enabled ON whitelist(enabled);
|
||||
|
||||
-- Insert default whitelist entries
|
||||
INSERT OR IGNORE INTO whitelist (symbol, name, wkn, isin, enabled, notes) VALUES
|
||||
('AAPL', 'Apple Inc.', '865985', 'US0378331005', TRUE, 'Technology - Consumer Electronics'),
|
||||
('MSFT', 'Microsoft Corporation', '870747', 'US5949181045', TRUE, 'Technology - Software'),
|
||||
('GOOGL', 'Alphabet Inc.', 'A14Y6F', 'US02079K3059', TRUE, 'Technology - Internet'),
|
||||
('TSLA', 'Tesla Inc.', 'A1CX3T', 'US88160R1014', TRUE, 'Automotive - Electric Vehicles'),
|
||||
('AMZN', 'Amazon.com Inc.', '906866', 'US0231351067', TRUE, 'E-Commerce - Cloud Computing');
|
||||
@@ -0,0 +1,18 @@
|
||||
-- Positions table to track open positions
|
||||
CREATE TABLE IF NOT EXISTS positions (
|
||||
id INTEGER PRIMARY KEY AUTOINCREMENT,
|
||||
symbol TEXT NOT NULL,
|
||||
quantity INTEGER NOT NULL,
|
||||
entry_price REAL NOT NULL,
|
||||
entry_trade_id INTEGER NOT NULL,
|
||||
current_price REAL,
|
||||
unrealized_pnl REAL,
|
||||
opened_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
updated_at TIMESTAMP DEFAULT CURRENT_TIMESTAMP,
|
||||
|
||||
FOREIGN KEY (entry_trade_id) REFERENCES trades(id),
|
||||
UNIQUE(symbol) -- Only one open position per symbol
|
||||
);
|
||||
|
||||
CREATE INDEX IF NOT EXISTS idx_positions_symbol ON positions(symbol);
|
||||
CREATE INDEX IF NOT EXISTS idx_positions_entry_trade ON positions(entry_trade_id);
|
||||
@@ -0,0 +1,10 @@
|
||||
-- Add LLM sentiment tracking columns
|
||||
|
||||
ALTER TABLE news_articles ADD COLUMN llm_sentiment_score REAL;
|
||||
ALTER TABLE news_articles ADD COLUMN llm_model TEXT;
|
||||
ALTER TABLE news_articles ADD COLUMN llm_confidence REAL;
|
||||
ALTER TABLE news_articles ADD COLUMN sentiment_method TEXT
|
||||
CHECK(sentiment_method IN ('keyword', 'llm', 'ensemble', 'keyword_fallback'));
|
||||
|
||||
-- Index for querying by sentiment method
|
||||
CREATE INDEX IF NOT EXISTS idx_news_sentiment_method ON news_articles(sentiment_method);
|
||||
+156
@@ -0,0 +1,156 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type NewsRepository struct {
|
||||
db Database
|
||||
}
|
||||
|
||||
func NewNewsRepository(db Database) *NewsRepository {
|
||||
return &NewsRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *NewsRepository) Create(ctx context.Context, article *model.NewsArticle) error {
|
||||
query := `
|
||||
INSERT INTO news_articles (
|
||||
source, title, url, content, published_at, fetched_at, symbols,
|
||||
sentiment_score, sentiment_label,
|
||||
llm_sentiment_score, llm_model, llm_confidence, sentiment_method
|
||||
)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
ON CONFLICT(url) DO NOTHING
|
||||
`
|
||||
|
||||
result, err := r.db.ExecContext(ctx, query,
|
||||
article.Source,
|
||||
article.Title,
|
||||
article.URL,
|
||||
article.Content,
|
||||
article.PublishedAt,
|
||||
article.FetchedAt,
|
||||
article.Symbols,
|
||||
article.SentimentScore,
|
||||
article.SentimentLabel,
|
||||
article.LLMSentimentScore,
|
||||
article.LLMModel,
|
||||
article.LLMConfidence,
|
||||
article.SentimentMethod,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert news article: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
// ON CONFLICT DO NOTHING means no rows affected, but not an error
|
||||
rowsAffected, _ := result.RowsAffected()
|
||||
if rowsAffected == 0 {
|
||||
return nil // Duplicate, silently ignore
|
||||
}
|
||||
return fmt.Errorf("failed to get last insert id: %w", err)
|
||||
}
|
||||
|
||||
article.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *NewsRepository) GetRecent(ctx context.Context, limit int) ([]*model.NewsArticle, error) {
|
||||
query := `
|
||||
SELECT id, source, title, url, content, published_at, fetched_at, symbols,
|
||||
sentiment_score, sentiment_label,
|
||||
llm_sentiment_score, llm_model, llm_confidence, sentiment_method
|
||||
FROM news_articles
|
||||
ORDER BY published_at DESC
|
||||
LIMIT ?
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query, limit)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query recent news: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var articles []*model.NewsArticle
|
||||
for rows.Next() {
|
||||
var article model.NewsArticle
|
||||
if err := rows.Scan(
|
||||
&article.ID,
|
||||
&article.Source,
|
||||
&article.Title,
|
||||
&article.URL,
|
||||
&article.Content,
|
||||
&article.PublishedAt,
|
||||
&article.FetchedAt,
|
||||
&article.Symbols,
|
||||
&article.SentimentScore,
|
||||
&article.SentimentLabel,
|
||||
&article.LLMSentimentScore,
|
||||
&article.LLMModel,
|
||||
&article.LLMConfidence,
|
||||
&article.SentimentMethod,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan news article: %w", err)
|
||||
}
|
||||
articles = append(articles, &article)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("rows error: %w", err)
|
||||
}
|
||||
|
||||
return articles, nil
|
||||
}
|
||||
|
||||
func (r *NewsRepository) GetBySymbol(ctx context.Context, symbol string, since time.Time, limit int) ([]*model.NewsArticle, error) {
|
||||
query := `
|
||||
SELECT id, source, title, url, content, published_at, fetched_at, symbols,
|
||||
sentiment_score, sentiment_label,
|
||||
llm_sentiment_score, llm_model, llm_confidence, sentiment_method
|
||||
FROM news_articles
|
||||
WHERE symbols LIKE ? AND published_at >= ?
|
||||
ORDER BY published_at DESC
|
||||
LIMIT ?
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query, "%"+symbol+"%", since, limit)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query news by symbol: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var articles []*model.NewsArticle
|
||||
for rows.Next() {
|
||||
var article model.NewsArticle
|
||||
if err := rows.Scan(
|
||||
&article.ID,
|
||||
&article.Source,
|
||||
&article.Title,
|
||||
&article.URL,
|
||||
&article.Content,
|
||||
&article.PublishedAt,
|
||||
&article.FetchedAt,
|
||||
&article.Symbols,
|
||||
&article.SentimentScore,
|
||||
&article.SentimentLabel,
|
||||
&article.LLMSentimentScore,
|
||||
&article.LLMModel,
|
||||
&article.LLMConfidence,
|
||||
&article.SentimentMethod,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan news article: %w", err)
|
||||
}
|
||||
articles = append(articles, &article)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating news rows: %w", err)
|
||||
}
|
||||
|
||||
return articles, nil
|
||||
}
|
||||
@@ -0,0 +1,154 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"database/sql"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type PositionRepository struct {
|
||||
db Database
|
||||
}
|
||||
|
||||
func NewPositionRepository(db Database) *PositionRepository {
|
||||
return &PositionRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *PositionRepository) Create(ctx context.Context, position *model.Position) error {
|
||||
query := `
|
||||
INSERT INTO positions (symbol, quantity, entry_price, entry_trade_id, opened_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?)
|
||||
`
|
||||
|
||||
now := time.Now()
|
||||
result, err := r.db.ExecContext(ctx, query,
|
||||
position.Symbol,
|
||||
position.Quantity,
|
||||
position.EntryPrice,
|
||||
position.EntryTradeID,
|
||||
now,
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert position: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get last insert id: %w", err)
|
||||
}
|
||||
|
||||
position.ID = id
|
||||
position.OpenedAt = now
|
||||
position.UpdatedAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PositionRepository) Update(ctx context.Context, position *model.Position) error {
|
||||
query := `
|
||||
UPDATE positions SET
|
||||
current_price = ?,
|
||||
unrealized_pnl = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
`
|
||||
|
||||
now := time.Now()
|
||||
_, err := r.db.ExecContext(ctx, query,
|
||||
position.CurrentPrice,
|
||||
position.UnrealizedPnL,
|
||||
now,
|
||||
position.ID,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update position: %w", err)
|
||||
}
|
||||
|
||||
position.UpdatedAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PositionRepository) Delete(ctx context.Context, id int64) error {
|
||||
query := `DELETE FROM positions WHERE id = ?`
|
||||
|
||||
_, err := r.db.ExecContext(ctx, query, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete position: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *PositionRepository) GetBySymbol(ctx context.Context, symbol string) (*model.Position, error) {
|
||||
query := `
|
||||
SELECT id, symbol, quantity, entry_price, entry_trade_id, current_price, unrealized_pnl, opened_at, updated_at
|
||||
FROM positions
|
||||
WHERE symbol = ?
|
||||
`
|
||||
|
||||
var position model.Position
|
||||
err := r.db.QueryRowContext(ctx, query, symbol).Scan(
|
||||
&position.ID,
|
||||
&position.Symbol,
|
||||
&position.Quantity,
|
||||
&position.EntryPrice,
|
||||
&position.EntryTradeID,
|
||||
&position.CurrentPrice,
|
||||
&position.UnrealizedPnL,
|
||||
&position.OpenedAt,
|
||||
&position.UpdatedAt,
|
||||
)
|
||||
|
||||
if err == sql.ErrNoRows {
|
||||
return nil, nil // No position found
|
||||
}
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get position: %w", err)
|
||||
}
|
||||
|
||||
return &position, nil
|
||||
}
|
||||
|
||||
func (r *PositionRepository) GetAll(ctx context.Context) ([]*model.Position, error) {
|
||||
query := `
|
||||
SELECT id, symbol, quantity, entry_price, entry_trade_id, current_price, unrealized_pnl, opened_at, updated_at
|
||||
FROM positions
|
||||
ORDER BY opened_at DESC
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query positions: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var positions []*model.Position
|
||||
for rows.Next() {
|
||||
var position model.Position
|
||||
if err := rows.Scan(
|
||||
&position.ID,
|
||||
&position.Symbol,
|
||||
&position.Quantity,
|
||||
&position.EntryPrice,
|
||||
&position.EntryTradeID,
|
||||
&position.CurrentPrice,
|
||||
&position.UnrealizedPnL,
|
||||
&position.OpenedAt,
|
||||
&position.UpdatedAt,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan position: %w", err)
|
||||
}
|
||||
positions = append(positions, &position)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating position rows: %w", err)
|
||||
}
|
||||
|
||||
return positions, nil
|
||||
}
|
||||
@@ -0,0 +1,253 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type TradeRepository struct {
|
||||
db Database
|
||||
}
|
||||
|
||||
func NewTradeRepository(db Database) *TradeRepository {
|
||||
return &TradeRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *TradeRepository) Create(ctx context.Context, trade *model.Trade) error {
|
||||
query := `
|
||||
INSERT INTO trades (
|
||||
symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, stop_loss_price, pending_until, created_at,
|
||||
is_dry_run
|
||||
) VALUES (?, ?, ?, ?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`
|
||||
|
||||
result, err := r.db.ExecContext(ctx, query,
|
||||
trade.Symbol,
|
||||
trade.Action,
|
||||
trade.Quantity,
|
||||
trade.Status,
|
||||
trade.Confidence,
|
||||
trade.Reasoning,
|
||||
trade.TargetPrice,
|
||||
trade.StopLossPrice,
|
||||
trade.PendingUtil,
|
||||
trade.CreatedAt,
|
||||
trade.IsDryRun,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert trade: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get last insert id: %w", err)
|
||||
}
|
||||
|
||||
trade.ID = id
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *TradeRepository) Update(ctx context.Context, trade *model.Trade) error {
|
||||
query := `
|
||||
UPDATE trades SET
|
||||
status = ?,
|
||||
executed_price = ?,
|
||||
ib_order_id = ?,
|
||||
approved_at = ?,
|
||||
rejected_at = ?,
|
||||
submitted_at = ?,
|
||||
filled_at = ?,
|
||||
completed_at = ?,
|
||||
rejection_reason = ?,
|
||||
forced_by_user = ?,
|
||||
dry_run_pnl = ?
|
||||
WHERE id = ?
|
||||
`
|
||||
|
||||
_, err := r.db.ExecContext(ctx, query,
|
||||
trade.Status,
|
||||
trade.ExecutedPrice,
|
||||
trade.IBOrderID,
|
||||
trade.ApprovedAt,
|
||||
trade.RejectedAt,
|
||||
trade.SubmittedAt,
|
||||
trade.FilledAt,
|
||||
trade.CompletedAt,
|
||||
trade.RejectionReason,
|
||||
trade.ForcedByUser,
|
||||
trade.DryRunPnL,
|
||||
trade.ID,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update trade: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *TradeRepository) GetByID(ctx context.Context, id int64) (*model.Trade, error) {
|
||||
query := `
|
||||
SELECT id, symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, executed_price, stop_loss_price, ib_order_id,
|
||||
created_at, pending_until, approved_at, rejected_at, submitted_at,
|
||||
filled_at, completed_at, rejection_reason, forced_by_user,
|
||||
is_dry_run, dry_run_pnl
|
||||
FROM trades
|
||||
WHERE id = ?
|
||||
`
|
||||
|
||||
var trade model.Trade
|
||||
err := r.db.QueryRowContext(ctx, query, id).Scan(
|
||||
&trade.ID,
|
||||
&trade.Symbol,
|
||||
&trade.Action,
|
||||
&trade.Quantity,
|
||||
&trade.Status,
|
||||
&trade.Confidence,
|
||||
&trade.Reasoning,
|
||||
&trade.TargetPrice,
|
||||
&trade.ExecutedPrice,
|
||||
&trade.StopLossPrice,
|
||||
&trade.IBOrderID,
|
||||
&trade.CreatedAt,
|
||||
&trade.PendingUtil,
|
||||
&trade.ApprovedAt,
|
||||
&trade.RejectedAt,
|
||||
&trade.SubmittedAt,
|
||||
&trade.FilledAt,
|
||||
&trade.CompletedAt,
|
||||
&trade.RejectionReason,
|
||||
&trade.ForcedByUser,
|
||||
&trade.IsDryRun,
|
||||
&trade.DryRunPnL,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get trade: %w", err)
|
||||
}
|
||||
|
||||
return &trade, nil
|
||||
}
|
||||
|
||||
func (r *TradeRepository) GetPendingTrades(ctx context.Context) ([]*model.Trade, error) {
|
||||
query := `
|
||||
SELECT id, symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, executed_price, stop_loss_price, ib_order_id,
|
||||
created_at, pending_until, approved_at, rejected_at, submitted_at,
|
||||
filled_at, completed_at, rejection_reason, forced_by_user,
|
||||
is_dry_run, dry_run_pnl
|
||||
FROM trades
|
||||
WHERE status = ?
|
||||
ORDER BY created_at ASC
|
||||
`
|
||||
|
||||
return r.queryTrades(ctx, query, model.TradePending)
|
||||
}
|
||||
|
||||
func (r *TradeRepository) GetExpiredPendingTrades(ctx context.Context) ([]*model.Trade, error) {
|
||||
query := `
|
||||
SELECT id, symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, executed_price, stop_loss_price, ib_order_id,
|
||||
created_at, pending_until, approved_at, rejected_at, submitted_at,
|
||||
filled_at, completed_at, rejection_reason, forced_by_user,
|
||||
is_dry_run, dry_run_pnl
|
||||
FROM trades
|
||||
WHERE status = ? AND pending_until <= ?
|
||||
ORDER BY created_at ASC
|
||||
`
|
||||
|
||||
return r.queryTrades(ctx, query, model.TradePending, time.Now())
|
||||
}
|
||||
|
||||
func (r *TradeRepository) GetActiveTrades(ctx context.Context) ([]*model.Trade, error) {
|
||||
query := `
|
||||
SELECT id, symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, executed_price, stop_loss_price, ib_order_id,
|
||||
created_at, pending_until, approved_at, rejected_at, submitted_at,
|
||||
filled_at, completed_at, rejection_reason, forced_by_user,
|
||||
is_dry_run, dry_run_pnl
|
||||
FROM trades
|
||||
WHERE status IN (?, ?, ?)
|
||||
ORDER BY created_at ASC
|
||||
`
|
||||
|
||||
return r.queryTrades(ctx, query, model.TradeSubmitted, model.TradeFilled, model.TradeCompleted)
|
||||
}
|
||||
|
||||
func (r *TradeRepository) GetTradesSince(ctx context.Context, since time.Time) ([]*model.Trade, error) {
|
||||
query := `
|
||||
SELECT id, symbol, action, quantity, status, confidence, reasoning,
|
||||
target_price, executed_price, stop_loss_price, ib_order_id,
|
||||
created_at, pending_until, approved_at, rejected_at, submitted_at,
|
||||
filled_at, completed_at, rejection_reason, forced_by_user,
|
||||
is_dry_run, dry_run_pnl
|
||||
FROM trades
|
||||
WHERE created_at >= ?
|
||||
ORDER BY created_at DESC
|
||||
`
|
||||
|
||||
return r.queryTrades(ctx, query, since)
|
||||
}
|
||||
|
||||
func (r *TradeRepository) CountTradesSince(ctx context.Context, since time.Time) (int, error) {
|
||||
query := `SELECT COUNT(*) FROM trades WHERE created_at >= ? AND status NOT IN (?, ?)`
|
||||
|
||||
var count int
|
||||
err := r.db.QueryRowContext(ctx, query, since, model.TradePending, model.TradeRejected).Scan(&count)
|
||||
if err != nil {
|
||||
return 0, fmt.Errorf("failed to count trades: %w", err)
|
||||
}
|
||||
|
||||
return count, nil
|
||||
}
|
||||
|
||||
func (r *TradeRepository) queryTrades(ctx context.Context, query string, args ...any) ([]*model.Trade, error) {
|
||||
rows, err := r.db.QueryContext(ctx, query, args...)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query trades: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var trades []*model.Trade
|
||||
for rows.Next() {
|
||||
var trade model.Trade
|
||||
if err := rows.Scan(
|
||||
&trade.ID,
|
||||
&trade.Symbol,
|
||||
&trade.Action,
|
||||
&trade.Quantity,
|
||||
&trade.Status,
|
||||
&trade.Confidence,
|
||||
&trade.Reasoning,
|
||||
&trade.TargetPrice,
|
||||
&trade.ExecutedPrice,
|
||||
&trade.StopLossPrice,
|
||||
&trade.IBOrderID,
|
||||
&trade.CreatedAt,
|
||||
&trade.PendingUtil,
|
||||
&trade.ApprovedAt,
|
||||
&trade.RejectedAt,
|
||||
&trade.SubmittedAt,
|
||||
&trade.FilledAt,
|
||||
&trade.CompletedAt,
|
||||
&trade.RejectionReason,
|
||||
&trade.ForcedByUser,
|
||||
&trade.IsDryRun,
|
||||
&trade.DryRunPnL,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan trade: %w", err)
|
||||
}
|
||||
trades = append(trades, &trade)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating trade rows: %w", err)
|
||||
}
|
||||
|
||||
return trades, nil
|
||||
}
|
||||
@@ -0,0 +1,236 @@
|
||||
package db
|
||||
|
||||
import (
|
||||
"context"
|
||||
"fmt"
|
||||
"time"
|
||||
|
||||
"github.com/pheinrich/aitrade/pkg/model"
|
||||
)
|
||||
|
||||
type WhitelistRepository struct {
|
||||
db Database
|
||||
}
|
||||
|
||||
func NewWhitelistRepository(db Database) *WhitelistRepository {
|
||||
return &WhitelistRepository{db: db}
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) Create(ctx context.Context, entry *model.WhitelistEntry) error {
|
||||
query := `
|
||||
INSERT INTO whitelist (symbol, name, wkn, isin, enabled, notes, created_at, updated_at)
|
||||
VALUES (?, ?, ?, ?, ?, ?, ?, ?)
|
||||
`
|
||||
|
||||
now := time.Now()
|
||||
result, err := r.db.ExecContext(ctx, query,
|
||||
entry.Symbol,
|
||||
entry.Name,
|
||||
entry.WKN,
|
||||
entry.ISIN,
|
||||
entry.Enabled,
|
||||
entry.Notes,
|
||||
now,
|
||||
now,
|
||||
)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to insert whitelist entry: %w", err)
|
||||
}
|
||||
|
||||
id, err := result.LastInsertId()
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to get last insert id: %w", err)
|
||||
}
|
||||
|
||||
entry.ID = id
|
||||
entry.CreatedAt = now
|
||||
entry.UpdatedAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) Update(ctx context.Context, entry *model.WhitelistEntry) error {
|
||||
query := `
|
||||
UPDATE whitelist SET
|
||||
name = ?,
|
||||
wkn = ?,
|
||||
isin = ?,
|
||||
enabled = ?,
|
||||
notes = ?,
|
||||
updated_at = ?
|
||||
WHERE id = ?
|
||||
`
|
||||
|
||||
now := time.Now()
|
||||
_, err := r.db.ExecContext(ctx, query,
|
||||
entry.Name,
|
||||
entry.WKN,
|
||||
entry.ISIN,
|
||||
entry.Enabled,
|
||||
entry.Notes,
|
||||
now,
|
||||
entry.ID,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to update whitelist entry: %w", err)
|
||||
}
|
||||
|
||||
entry.UpdatedAt = now
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) Delete(ctx context.Context, id int64) error {
|
||||
query := `DELETE FROM whitelist WHERE id = ?`
|
||||
|
||||
_, err := r.db.ExecContext(ctx, query, id)
|
||||
if err != nil {
|
||||
return fmt.Errorf("failed to delete whitelist entry: %w", err)
|
||||
}
|
||||
|
||||
return nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) GetByID(ctx context.Context, id int64) (*model.WhitelistEntry, error) {
|
||||
query := `
|
||||
SELECT id, symbol, name, wkn, isin, enabled, created_at, updated_at, notes
|
||||
FROM whitelist
|
||||
WHERE id = ?
|
||||
`
|
||||
|
||||
var entry model.WhitelistEntry
|
||||
err := r.db.QueryRowContext(ctx, query, id).Scan(
|
||||
&entry.ID,
|
||||
&entry.Symbol,
|
||||
&entry.Name,
|
||||
&entry.WKN,
|
||||
&entry.ISIN,
|
||||
&entry.Enabled,
|
||||
&entry.CreatedAt,
|
||||
&entry.UpdatedAt,
|
||||
&entry.Notes,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get whitelist entry: %w", err)
|
||||
}
|
||||
|
||||
return &entry, nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) GetBySymbol(ctx context.Context, symbol string) (*model.WhitelistEntry, error) {
|
||||
query := `
|
||||
SELECT id, symbol, name, wkn, isin, enabled, created_at, updated_at, notes
|
||||
FROM whitelist
|
||||
WHERE symbol = ?
|
||||
`
|
||||
|
||||
var entry model.WhitelistEntry
|
||||
err := r.db.QueryRowContext(ctx, query, symbol).Scan(
|
||||
&entry.ID,
|
||||
&entry.Symbol,
|
||||
&entry.Name,
|
||||
&entry.WKN,
|
||||
&entry.ISIN,
|
||||
&entry.Enabled,
|
||||
&entry.CreatedAt,
|
||||
&entry.UpdatedAt,
|
||||
&entry.Notes,
|
||||
)
|
||||
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to get whitelist entry: %w", err)
|
||||
}
|
||||
|
||||
return &entry, nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) GetAll(ctx context.Context) ([]*model.WhitelistEntry, error) {
|
||||
query := `
|
||||
SELECT id, symbol, name, wkn, isin, enabled, created_at, updated_at, notes
|
||||
FROM whitelist
|
||||
ORDER BY symbol ASC
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query whitelist: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var entries []*model.WhitelistEntry
|
||||
for rows.Next() {
|
||||
var entry model.WhitelistEntry
|
||||
if err := rows.Scan(
|
||||
&entry.ID,
|
||||
&entry.Symbol,
|
||||
&entry.Name,
|
||||
&entry.WKN,
|
||||
&entry.ISIN,
|
||||
&entry.Enabled,
|
||||
&entry.CreatedAt,
|
||||
&entry.UpdatedAt,
|
||||
&entry.Notes,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan whitelist entry: %w", err)
|
||||
}
|
||||
entries = append(entries, &entry)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating whitelist rows: %w", err)
|
||||
}
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) GetEnabled(ctx context.Context) ([]*model.WhitelistEntry, error) {
|
||||
query := `
|
||||
SELECT id, symbol, name, wkn, isin, enabled, created_at, updated_at, notes
|
||||
FROM whitelist
|
||||
WHERE enabled = TRUE
|
||||
ORDER BY symbol ASC
|
||||
`
|
||||
|
||||
rows, err := r.db.QueryContext(ctx, query)
|
||||
if err != nil {
|
||||
return nil, fmt.Errorf("failed to query enabled whitelist: %w", err)
|
||||
}
|
||||
defer rows.Close()
|
||||
|
||||
var entries []*model.WhitelistEntry
|
||||
for rows.Next() {
|
||||
var entry model.WhitelistEntry
|
||||
if err := rows.Scan(
|
||||
&entry.ID,
|
||||
&entry.Symbol,
|
||||
&entry.Name,
|
||||
&entry.WKN,
|
||||
&entry.ISIN,
|
||||
&entry.Enabled,
|
||||
&entry.CreatedAt,
|
||||
&entry.UpdatedAt,
|
||||
&entry.Notes,
|
||||
); err != nil {
|
||||
return nil, fmt.Errorf("failed to scan whitelist entry: %w", err)
|
||||
}
|
||||
entries = append(entries, &entry)
|
||||
}
|
||||
|
||||
if err := rows.Err(); err != nil {
|
||||
return nil, fmt.Errorf("error iterating whitelist rows: %w", err)
|
||||
}
|
||||
|
||||
return entries, nil
|
||||
}
|
||||
|
||||
func (r *WhitelistRepository) IsSymbolWhitelisted(ctx context.Context, symbol string) (bool, error) {
|
||||
query := `SELECT COUNT(*) FROM whitelist WHERE symbol = ? AND enabled = TRUE`
|
||||
|
||||
var count int
|
||||
err := r.db.QueryRowContext(ctx, query, symbol).Scan(&count)
|
||||
if err != nil {
|
||||
return false, fmt.Errorf("failed to check whitelist: %w", err)
|
||||
}
|
||||
|
||||
return count > 0, nil
|
||||
}
|
||||
@@ -0,0 +1,13 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type Balance struct {
|
||||
ID int64
|
||||
Timestamp time.Time
|
||||
TotalValue float64
|
||||
CashBalance float64
|
||||
BuyingPower float64
|
||||
UnrealizedPnL *float64
|
||||
RealizedPnL *float64
|
||||
}
|
||||
@@ -0,0 +1,21 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type NewsArticle struct {
|
||||
ID int64
|
||||
Source string
|
||||
Title string
|
||||
URL string
|
||||
Content string
|
||||
PublishedAt time.Time
|
||||
FetchedAt time.Time
|
||||
Symbols string
|
||||
SentimentScore *float64
|
||||
SentimentLabel string
|
||||
// LLM sentiment fields
|
||||
LLMSentimentScore *float64
|
||||
LLMModel *string
|
||||
LLMConfidence *float64
|
||||
SentimentMethod string // "keyword", "llm", "ensemble", "keyword_fallback"
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type Position struct {
|
||||
ID int64
|
||||
Symbol string
|
||||
Quantity int
|
||||
EntryPrice float64
|
||||
EntryTradeID int64
|
||||
CurrentPrice *float64
|
||||
UnrealizedPnL *float64
|
||||
OpenedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
}
|
||||
@@ -0,0 +1,52 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type TradeStatus string
|
||||
|
||||
const (
|
||||
TradePending TradeStatus = "PENDING"
|
||||
TradeApproved TradeStatus = "APPROVED"
|
||||
TradeRejected TradeStatus = "REJECTED"
|
||||
TradeSubmitted TradeStatus = "SUBMITTED"
|
||||
TradeFilled TradeStatus = "FILLED"
|
||||
TradeCompleted TradeStatus = "COMPLETED"
|
||||
TradeStopped TradeStatus = "STOPPED"
|
||||
)
|
||||
|
||||
type ActionType string
|
||||
|
||||
const (
|
||||
ActionBuy ActionType = "BUY"
|
||||
ActionSell ActionType = "SELL"
|
||||
)
|
||||
|
||||
type Trade struct {
|
||||
ID int64
|
||||
Symbol string
|
||||
Action ActionType
|
||||
Quantity int
|
||||
Status TradeStatus
|
||||
Confidence float64
|
||||
Reasoning string
|
||||
|
||||
TargetPrice *float64
|
||||
ExecutedPrice *float64
|
||||
StopLossPrice *float64
|
||||
IBOrderID *int64
|
||||
|
||||
CreatedAt time.Time
|
||||
PendingUtil *time.Time
|
||||
ApprovedAt *time.Time
|
||||
RejectedAt *time.Time
|
||||
SubmittedAt *time.Time
|
||||
FilledAt *time.Time
|
||||
CompletedAt *time.Time
|
||||
|
||||
RejectionReason *string
|
||||
ForcedByUser bool
|
||||
|
||||
// Dry run fields
|
||||
IsDryRun bool
|
||||
DryRunPnL *float64 // P&L for this trade in dry run
|
||||
}
|
||||
@@ -0,0 +1,15 @@
|
||||
package model
|
||||
|
||||
import "time"
|
||||
|
||||
type WhitelistEntry struct {
|
||||
ID int64
|
||||
Symbol string
|
||||
Name string
|
||||
WKN string
|
||||
ISIN string
|
||||
Enabled bool
|
||||
CreatedAt time.Time
|
||||
UpdatedAt time.Time
|
||||
Notes string
|
||||
}
|
||||
Reference in new issue
Block a user