import requests
import time
import concurrent.futures
import statistics

def test_single_request(prompt, model="qwen3:0.6b"):
    """测试单个请求的响应时间"""
    url = "http://localhost:11434/api/generate"
    data = {
        "model": model,
        "prompt": prompt,
        "stream": False
    }
    
    start_time = time.time()
    response = requests.post(url, json=data)
    end_time = time.time()
    
    if response.status_code == 200:
        return end_time - start_time, len(response.json()["response"])
    else:
        return None, 0

def test_concurrent_requests(num_requests=10, max_workers=5):
    """测试并发请求"""
    prompts = ["请介绍一下人工智能的发展历史"] * num_requests
    
    with concurrent.futures.ThreadPoolExecutor(max_workers=max_workers) as executor:
        start_time = time.time()
        futures = [executor.submit(test_single_request, prompt) for prompt in prompts]
        results = [future.result() for future in futures]
        end_time = time.time()
    
    # 过滤成功的结果
    successful_results = [r for r in results if r[0] is not None]
    response_times = [r[0] for r in successful_results]
    token_counts = [r[1] for r in successful_results]
    
    print(f"总请求数: {num_requests}")
    print(f"成功请求数: {len(successful_results)}")
    print(f"总耗时: {end_time - start_time:.2f}秒")
    print(f"平均响应时间: {statistics.mean(response_times):.2f}秒")
    print(f"响应时间中位数: {statistics.median(response_times):.2f}秒")
    print(f"平均生成token数: {statistics.mean(token_counts):.1f}")
    print(f"吞吐量: {len(successful_results)/(end_time - start_time):.2f} 请求/秒")

if __name__ == "__main__":
    # 测试单个请求
    print("=== 单个请求测试 ===")
    response_time, token_count = test_single_request("请介绍一下人工智能的发展历史")
    if response_time:
        print(f"响应时间: {response_time:.2f}秒")
        print(f"生成token数: {token_count}")
        print(f"生成速度: {token_count/response_time:.1f} tokens/秒")
    
    # 测试并发请求
    print("\n=== 并发请求测试 ===")
    test_concurrent_requests(num_requests=20, max_workers=5)
