Files
stock/stock-html/routes/smart_trade.py
T

288 lines
11 KiB
Python

"""
智能交易引擎 API 路由
提供算法配置管理、引擎状态查看、手动触发等功能
"""
from flask import Blueprint, request, jsonify
from datetime import date
from db import get_db, login_required, get_current_user_id
from psycopg2.extras import RealDictCursor
bp = Blueprint('smart_trade', __name__, url_prefix='/api/smart')
# ═══════════════════════════════════════════════════════
# 1. 算法模板
# ═══════════════════════════════════════════════════════
@bp.route('/templates', methods=['GET'])
@login_required
def get_algo_templates():
"""获取所有预置算法模板"""
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
cur = conn.cursor(cursor_factory=RealDictCursor)
cur.execute("""
SELECT id, name, display_name, description, risk_level,
take_profit_pct::float, stop_loss_pct::float,
ignore_sell_signal, sell_confirm_days,
max_hold_days, no_timeout_if_rising,
position_pct::float, signal_weight,
partial_exit_pct, momentum_trail_gap::float,
breakeven_at::float, momentum_tp, momentum_days,
buy_time, sell_time,
backtest_annual_return::float, backtest_max_drawdown::float,
backtest_win_rate::float, backtest_calmar::float
FROM algo_templates
ORDER BY backtest_calmar DESC NULLS LAST
""")
templates = cur.fetchall()
return jsonify({'success': True, 'templates': templates})
except Exception as e:
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
# ═══════════════════════════════════════════════════════
# 2. 用户算法配置 (CRUD)
# ═══════════════════════════════════════════════════════
@bp.route('/config', methods=['GET'])
@login_required
def get_algo_config():
"""获取用户当前的算法配置"""
user_id = get_current_user_id()
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import get_user_algo_config, DEFAULT_CONFIG
config = get_user_algo_config(conn, user_id)
# 判断是否是默认配置(没有存入数据库)
cur = conn.cursor(cursor_factory=RealDictCursor)
cur.execute("SELECT COUNT(*) as cnt FROM sim_algo_config WHERE user_id = %s", (user_id,))
has_config = cur.fetchone()['cnt'] > 0
return jsonify({
'success': True,
'config': config,
'is_default': not has_config,
})
except Exception as e:
import traceback
traceback.print_exc()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
@bp.route('/config', methods=['POST'])
@login_required
def save_algo_config():
"""保存/更新用户的算法配置"""
user_id = get_current_user_id()
data = request.get_json()
if not data:
return jsonify({'success': False, 'error': '无效参数'}), 400
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import save_user_algo_config, DEFAULT_CONFIG
# 合并默认值
config = dict(DEFAULT_CONFIG)
for key in config:
if key in data:
config[key] = data[key]
save_user_algo_config(conn, user_id, config)
return jsonify({'success': True, 'message': '算法配置已保存'})
except Exception as e:
conn.rollback()
import traceback
traceback.print_exc()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
@bp.route('/apply_template', methods=['POST'])
@login_required
def apply_template():
"""从模板应用算法配置"""
user_id = get_current_user_id()
data = request.get_json()
template_name = data.get('template_name')
if not template_name:
return jsonify({'success': False, 'error': '缺少template_name'}), 400
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import apply_template as do_apply
ok = do_apply(conn, user_id, template_name)
if ok:
return jsonify({'success': True, 'message': f'已应用模板: {template_name}'})
else:
return jsonify({'success': False, 'error': f'模板不存在: {template_name}'}), 404
except Exception as e:
conn.rollback()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
# ═══════════════════════════════════════════════════════
# 3. 引擎状态
# ═══════════════════════════════════════════════════════
@bp.route('/status', methods=['GET'])
@login_required
def get_status():
"""获取智能交易引擎的当前状态(含持仓详情+活跃规则+信号日志)"""
user_id = get_current_user_id()
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import get_engine_status
status = get_engine_status(conn, user_id)
return jsonify({'success': True, **status})
except Exception as e:
import traceback
traceback.print_exc()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
# ═══════════════════════════════════════════════════════
# 4. 手动触发
# ═══════════════════════════════════════════════════════
@bp.route('/trigger', methods=['POST'])
@login_required
def trigger_smart_trade():
"""手动触发智能交易引擎执行"""
user_id = get_current_user_id()
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import execute_smart_trade
result = execute_smart_trade(conn, user_id, scan_date=None)
if result.get('error') and not result.get('success'):
return jsonify({'success': False, 'error': result['error']}), 500
return jsonify({
'success': True,
'results': result.get('results', []),
'signals': result.get('signals', 0),
'algo': result.get('algo', 'unknown'),
'total_fees': result.get('total_fees', 0),
'detail_reasons': result.get('detail_reasons', []),
'skipped_limit': result.get('skipped_limit', []),
'skipped_t1': result.get('skipped_t1', []),
'available_cash': result.get('available_cash', 0),
'message': f"智能引擎执行完成: {result.get('signals', 0)}笔信号"
})
except Exception as e:
import traceback
traceback.print_exc()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
# ═══════════════════════════════════════════════════════
# 5. 信号日志
# ═══════════════════════════════════════════════════════
@bp.route('/signals', methods=['GET'])
@login_required
def get_signals():
"""获取交易信号日志"""
user_id = get_current_user_id()
limit = request.args.get('limit', 50, type=int)
days = request.args.get('days', 7, type=int)
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
cur = conn.cursor(cursor_factory=RealDictCursor)
cur.execute("""
SELECT id, signal_date::text, signal_time::text,
stock_code, stock_name, action, reason, algo_rule,
signal_price::float, buy_price::float, profit_pct::float,
executed, execute_price::float, execute_shares,
created_at::text
FROM sim_trade_signals
WHERE user_id = %s AND signal_date >= CURRENT_DATE - %s
ORDER BY signal_date DESC, id DESC
LIMIT %s
""", (user_id, days, limit))
signals = cur.fetchall()
return jsonify({'success': True, 'signals': signals})
except Exception as e:
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)
# ═══════════════════════════════════════════════════════
# 6. 持仓元数据(前端持仓详情扩展)
# ═══════════════════════════════════════════════════════
@bp.route('/position_meta', methods=['GET'])
@login_required
def get_position_meta():
"""获取持仓的算法元数据(止盈止损状态等)"""
user_id = get_current_user_id()
conn = get_db()
if not conn:
return jsonify({'success': False, 'error': '数据库连接失败'}), 500
try:
from services.smart_trade_engine import get_all_position_meta
positions = get_all_position_meta(conn, user_id)
result = []
for p in positions:
result.append({
'stock_code': p['stock_code'],
'buy_date': str(p.get('buy_date', '')),
'buy_price': float(p.get('buy_price', 0)),
'days_held': p.get('days_held', 0),
'max_price': float(p.get('max_price_since_buy', 0) or 0),
'consecutive_up_days': p.get('consecutive_up_days', 0),
'consecutive_sell_signals': p.get('consecutive_sell_signals', 0),
'partial_exit_done': p.get('partial_exit_done', False),
'breakeven_active': p.get('breakeven_active', False),
'momentum_trailing_active': p.get('momentum_trailing_active', False),
'momentum_high_price': float(p.get('momentum_high_price', 0) or 0),
'current_shares': p.get('current_shares', 0),
'original_shares': p.get('original_shares', 0),
})
return jsonify({'success': True, 'positions': result})
except Exception as e:
import traceback
traceback.print_exc()
return jsonify({'success': False, 'error': str(e)}), 500
finally:
put_db(conn)