# source: https://raw.githubusercontent.com/CETANGZHI/flowainew/94cc7095db1259c5f1cb053bdf5c74bb0c4c16c3/.archive/deprecated_scripts_20251219/backend_tests/test_api.py
"""
API测试脚本
用于测试后端API是否正常工作
"""

import requests
import json

BASE_URL = "http://localhost:8000"

# 全局变量存储token
AUTH_TOKEN = None


def test_health_check():
    """测试健康检查"""
    print("\n=== 测试健康检查 ===")
    response = requests.get(f"{BASE_URL}/health")
    print(f"状态码: {response.status_code}")
    print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    return response.status_code == 200


def test_root():
    """测试根路径"""
    print("\n=== 测试根路径 ===")
    response = requests.get(f"{BASE_URL}/")
    print(f"状态码: {response.status_code}")
    print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    return response.status_code == 200


def test_chat():
    """测试对话API"""
    print("\n=== 测试对话API ===")
    
    data = {
        "messages": [
            {"role": "user", "content": "你好，请介绍一下你自己"}
        ],
        "temperature": 0.7
    }
    
    response = requests.post(
        f"{BASE_URL}/v1/chat/",
        json=data,
        headers={"Content-Type": "application/json"}
    )
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")
    
    return response.status_code == 200


def test_chat_stream():
    """测试流式对话API"""
    print("\n=== 测试流式对话API ===")
    
    data = {
        "messages": [
            {"role": "user", "content": "请简单介绍一下比特币"}
        ],
        "temperature": 0.7
    }
    
    try:
        response = requests.post(
            f"{BASE_URL}/v1/chat/stream",
            json=data,
            headers={"Content-Type": "application/json"},
            stream=True
        )
        
        print(f"状态码: {response.status_code}")
        
        if response.status_code == 200:
            print("流式响应:")
            for line in response.iter_lines():
                if line:
                    line_str = line.decode('utf-8')
                    if line_str.startswith('data: '):
                        data_str = line_str[6:]
                        if data_str != '[DONE]':
                            try:
                                data_json = json.loads(data_str)
                                if data_json.get('type') == 'text-delta':
                                    print(data_json.get('textDelta', ''), end='', flush=True)
                            except json.JSONDecodeError:
                                pass
            print()  # 换行
            return True
        else:
            print(f"错误: {response.text}")
            return False
    
    except Exception as e:
        print(f"错误: {e}")
        return False


def test_create_chat():
    """测试创建对话"""
    print("\n=== 测试创建对话 ===")
    
    data = {
        "title": "测试对话"
    }
    
    response = requests.post(
        f"{BASE_URL}/v1/chat/chats",
        json=data,
        headers={"Content-Type": "application/json"}
    )
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        result = response.json()
        print(f"响应: {json.dumps(result, indent=2, ensure_ascii=False)}")
        return result.get('id')
    else:
        print(f"错误: {response.text}")
        return None


def test_get_chats():
    """测试获取对话列表"""
    print("\n=== 测试获取对话列表 ===")
    
    response = requests.get(f"{BASE_URL}/v1/chat/chats")
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")
    
    return response.status_code == 200


def test_backtest():
    """测试回测API"""
    print("\n=== 测试回测API ===")
    
    # 简单的策略代码
    strategy_code = """
from freqtrade.strategy import IStrategy
from pandas import DataFrame

class Github_CETANGZHI_flowainew__test_api__20260114_114646(IStrategy):
    minimal_roi = {"0": 0.10}
    stoploss = -0.05
    timeframe = '1h'
    
    def populate_indicators(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        return dataframe
    
    def populate_entry_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe.loc[:, 'enter_long'] = 1
        return dataframe
    
    def populate_exit_trend(self, dataframe: DataFrame, metadata: dict) -> DataFrame:
        dataframe.loc[:, 'exit_long'] = 1
        return dataframe
"""
    
    data = {
        "strategy_code": strategy_code,
        "pairs": ["BTC/USDT", "ETH/USDT"],
        "timerange": "20230101-20231231",
        "timeframe": "1h",
        "initial_balance": 10000
    }
    
    response = requests.post(
        f"{BASE_URL}/v1/backtest/run",
        json=data,
        headers={"Content-Type": "application/json"}
    )
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        result = response.json()
        print(f"响应: {json.dumps(result, indent=2, ensure_ascii=False)}")
        return result.get('task_id')
    else:
        print(f"错误: {response.text}")
        return None


def test_get_backtest_result(task_id: str):
    """测试获取回测结果"""
    print(f"\n=== 测试获取回测结果 (task_id: {task_id}) ===")
    
    response = requests.get(f"{BASE_URL}/v1/backtest/{task_id}")
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")
    
    return response.status_code == 200


def test_get_backtest_list():
    """测试获取回测列表"""
    print("\n=== 测试获取回测列表 ===")
    
    response = requests.get(f"{BASE_URL}/v1/backtest/")
    
    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")
    
    return response.status_code == 200


def test_market_ticker():
    """测试获取价格信息"""
    print("\n=== 测试获取价格信息 ===")

    response = requests.get(f"{BASE_URL}/v1/market/ticker?symbol=BTC/USDT")

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def test_market_ohlcv():
    """测试获取K线数据"""
    print("\n=== 测试获取K线数据 ===")

    response = requests.get(
        f"{BASE_URL}/v1/market/ohlcv?symbol=BTC/USDT&timeframe=1h&limit=10"
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        data = response.json()
        print(f"返回数量: {len(data)}")
        if data:
            print(f"最新K线: {json.dumps(data[-1], indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def test_market_indicators():
    """测试获取技术指标"""
    print("\n=== 测试获取技术指标 ===")

    response = requests.get(
        f"{BASE_URL}/v1/market/indicators?symbol=BTC/USDT&timeframe=1h&indicators=rsi,macd"
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def test_market_sentiment():
    """测试市场情绪分析"""
    print("\n=== 测试市场情绪分析 ===")

    response = requests.get(f"{BASE_URL}/v1/market/sentiment?symbol=BTC/USDT")

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def test_create_strategy():
    """测试创建策略"""
    print("\n=== 测试创建策略 ===")

    data = {
        "name": "测试策略",
        "description": "这是一个测试策略",
        "code": "# 策略代码",
        "pairs": ["BTC/USDT"],
        "timeframe": "1h",
        "is_public": False
    }

    response = requests.post(
        f"{BASE_URL}/v1/strategy/",
        json=data,
        headers={"Content-Type": "application/json"}
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        result = response.json()
        print(f"响应: {json.dumps(result, indent=2, ensure_ascii=False)}")
        return result.get('id')
    else:
        print(f"错误: {response.text}")
        return None


def test_get_strategies():
    """测试获取策略列表"""
    print("\n=== 测试获取策略列表 ===")

    response = requests.get(f"{BASE_URL}/v1/strategy/")

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def test_register():
    """测试用户注册"""
    print("\n=== 测试用户注册 ===")

    import random
    email = f"test{random.randint(1000, 9999)}@example.com"

    data = {
        "email": email,
        "password": "test123456",
        "name": "测试用户"
    }

    response = requests.post(
        f"{BASE_URL}/v1/auth/register",
        json=data,
        headers={"Content-Type": "application/json"}
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
        return email
    else:
        print(f"错误: {response.text}")
        return None


def test_login(email: str = "test@example.com", password: str = "test123456"):
    """测试用户登录"""
    print(f"\n=== 测试用户登录 ({email}) ===")

    global AUTH_TOKEN

    data = {
        "email": email,
        "password": password
    }

    response = requests.post(
        f"{BASE_URL}/v1/auth/login",
        json=data,
        headers={"Content-Type": "application/json"}
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        result = response.json()
        AUTH_TOKEN = result.get("access_token")
        print(f"Token: {AUTH_TOKEN[:50]}...")
        print(f"用户: {result.get('user', {}).get('email')}")
        return True
    else:
        print(f"错误: {response.text}")
        return False


def test_get_me():
    """测试获取当前用户信息"""
    print("\n=== 测试获取当前用户信息 ===")

    if not AUTH_TOKEN:
        print("⚠️  未登录，跳过测试")
        return False

    response = requests.get(
        f"{BASE_URL}/v1/auth/me",
        headers={"Authorization": f"Bearer {AUTH_TOKEN}"}
    )

    print(f"状态码: {response.status_code}")
    if response.status_code == 200:
        print(f"响应: {json.dumps(response.json(), indent=2, ensure_ascii=False)}")
    else:
        print(f"错误: {response.text}")

    return response.status_code == 200


def main():
    """运行所有测试"""
    print("=" * 60)
    print("开始测试后端API")
    print("=" * 60)

    results = {}

    # 基础测试
    results['health'] = test_health_check()
    results['root'] = test_root()

    # Chat API测试
    results['chat'] = test_chat()
    results['chat_stream'] = test_chat_stream()
    results['create_chat'] = test_create_chat() is not None
    results['get_chats'] = test_get_chats()

    # Backtest API测试
    task_id = test_backtest()
    results['backtest'] = task_id is not None

    if task_id:
        results['get_backtest_result'] = test_get_backtest_result(task_id)

    results['get_backtest_list'] = test_get_backtest_list()

    # Market API测试
    results['market_ticker'] = test_market_ticker()
    results['market_ohlcv'] = test_market_ohlcv()
    results['market_indicators'] = test_market_indicators()
    results['market_sentiment'] = test_market_sentiment()

    # Strategy API测试
    strategy_id = test_create_strategy()
    results['create_strategy'] = strategy_id is not None
    results['get_strategies'] = test_get_strategies()

    # Auth API测试
    new_email = test_register()
    results['register'] = new_email is not None

    # 使用测试用户登录
    results['login'] = test_login()
    results['get_me'] = test_get_me()
    
    # 打印测试结果
    print("\n" + "=" * 60)
    print("测试结果汇总")
    print("=" * 60)
    
    for test_name, passed in results.items():
        status = "✅ 通过" if passed else "❌ 失败"
        print(f"{test_name:30s} {status}")
    
    total = len(results)
    passed = sum(results.values())
    print(f"\n总计: {passed}/{total} 通过")
    
    return passed == total


if __name__ == "__main__":
    try:
        success = main()
        exit(0 if success else 1)
    except KeyboardInterrupt:
        print("\n\n测试被中断")
        exit(1)
    except Exception as e:
        print(f"\n\n测试出错: {e}")
        exit(1)

