SPB Git forge

spb/doc-api

Public
2commits 1branches 0releases
15.7 MBsize
maindefault branch
13 days agolast push
Python 88.3% TypeScript 7.6% Shell 4.1%
6.6 KB · 103 lines python
Raw Blame History
1"""xAI tools: function calling (cheap) and server-side tools (web_search, x_search, code_interpreter, mcp, file_search — RUN_EXPENSIVE_TESTS)."""2from __future__ import annotations34import json56import pytest78TOOLS_CHAT = [{"type": "function", "function": {"name": "get_weather", "description": "Get weather for a city", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}}}]9TOOLS_RESP = [{"type": "function", "name": "get_weather", "description": "Get weather for a city", "parameters": {"type": "object", "properties": {"city": {"type": "string"}}, "required": ["city"]}}]101112def _text(body):13    return next((c["text"] for o in body["output"] if o["type"] == "message" for c in o["content"] if c["type"] == "output_text"), None)141516def test_chat_function_call_roundtrip(xai, models):17    msgs = [{"role": "user", "content": "Weather in Paris?"}]18    st, b, _ = xai("POST", "/v1/chat/completions", {"model": models["xai"], "messages": msgs, "max_completion_tokens": 64, "tools": TOOLS_CHAT, "tool_choice": {"type": "function", "function": {"name": "get_weather"}}}, est_cost_usd=0.0005, note="test_tools chat forced")19    assert st == 200 and b["choices"][0]["finish_reason"] == "tool_calls"20    tc = b["choices"][0]["message"]["tool_calls"][0]21    assert tc["type"] == "function" and tc["function"]["name"] == "get_weather" and "city" in json.loads(tc["function"]["arguments"])22    msgs += [b["choices"][0]["message"], {"role": "tool", "tool_call_id": tc["id"], "content": "{\"temp_c\": 18}"}]23    st, b2, _ = xai("POST", "/v1/chat/completions", {"model": models["xai"], "messages": msgs, "max_completion_tokens": 32, "tools": TOOLS_CHAT}, est_cost_usd=0.0005, note="test_tools chat round trip")24    assert st == 200 and "18" in b2["choices"][0]["message"]["content"]252627def test_responses_function_call_roundtrip(xai, models):28    st, r, _ = xai("POST", "/v1/responses", {"model": models["xai"], "input": "Weather in Paris?", "tools": TOOLS_RESP, "tool_choice": {"type": "function", "name": "get_weather"}, "max_output_tokens": 2000}, est_cost_usd=0.0005, note="test_tools responses forced")29    assert st == 20030    fc = next(o for o in r["output"] if o["type"] == "function_call")31    assert fc["name"] == "get_weather" and fc["call_id"].startswith("call-")32    st, r2, _ = xai("POST", "/v1/responses", {"model": models["xai"], "previous_response_id": r["id"], "input": [{"type": "function_call_output", "call_id": fc["call_id"], "output": "{\"temp_c\": 18}"}], "tools": TOOLS_RESP, "max_output_tokens": 2000}, est_cost_usd=0.0005, note="test_tools responses function_call_output")33    assert st == 200 and "18" in _text(r2)34    for rid in (r["id"], r2["id"]):35        xai("DELETE", f"/v1/responses/{rid}", note="test_tools cleanup")363738def test_chat_rejects_server_tools(xai, models):39    st, b, _ = xai("POST", "/v1/chat/completions", {"model": models["xai"], "messages": [{"role": "user", "content": "Reply with OK."}], "max_completion_tokens": 32, "tools": [{"type": "web_search"}]}, est_cost_usd=0, note="test_tools chat web_search 422")40    assert st == 422 and "unknown variant" in (b.decode() if isinstance(b, bytes) else json.dumps(b))414243def test_unknown_tool_type_lists_accepted_types(xai, models):44    st, b, _ = xai("POST", "/v1/responses", {"model": models["xai"], "input": "Reply with OK.", "store": False, "tools": [{"type": "web_search_preview"}]}, est_cost_usd=0, note="test_tools web_search_preview 422")45    txt = b.decode() if isinstance(b, bytes) else json.dumps(b)46    assert st == 42247    for t in ("function", "web_search", "x_search", "code_interpreter", "file_search", "mcp"):48        assert f"`{t}`" in txt495051def _run_tool(xai, models, tools, prompt, note, include=None):52    body = {"model": models["xai"], "input": prompt, "tools": tools, "store": False, "max_output_tokens": 4000, "max_turns": 3}53    if include:54        body["include"] = include55    st, r, _ = xai("POST", "/v1/responses", body, est_cost_usd=0.02, note=note, timeout=300)56    assert st == 200, r57    return r585960@pytest.mark.run_expensive_tests61def test_web_search_tool(xai, models):62    r = _run_tool(xai, models, [{"type": "web_search"}], "Use the web_search tool to open https://x.ai/news and give the newest article title. Max 8 words.", "test_tools web_search", ["web_search_call.action.sources"])63    calls = [o for o in r["output"] if o["type"] == "web_search_call"]64    assert calls and calls[0]["action"]["type"] in ("search", "open_page", "find_in_page")65    assert r["usage"]["server_side_tool_usage_details"]["web_search_calls"] >= 1666768@pytest.mark.run_expensive_tests69def test_x_search_tool(xai, models):70    r = _run_tool(xai, models, [{"type": "x_search", "allowed_x_handles": ["xai"]}], "What is the latest post by @xai about? Max 8 words.", "test_tools x_search")71    types = {o["type"] for o in r["output"]}72    assert types & {"custom_tool_call", "x_search_call"}  # live: custom_tool_call (x_keyword_search…)73    assert r["usage"]["server_side_tool_usage_details"]["x_search_calls"] >= 1747576@pytest.mark.run_expensive_tests77def test_code_interpreter_tool(xai, models):78    r = _run_tool(xai, models, [{"type": "code_interpreter"}], "Run print(2+2) in Python and reply with the output only.", "test_tools code_interpreter", ["code_interpreter_call.outputs"])79    ci = next(o for o in r["output"] if o["type"] == "code_interpreter_call")80    assert "print" in ci["code"] and ci["outputs"] and ci["outputs"][0]["type"] == "logs"81    assert "4" in _text(r)828384@pytest.mark.run_expensive_tests85def test_mcp_tool(xai, models):86    r = _run_tool(xai, models, [{"type": "mcp", "server_url": "https://mcp.deepwiki.com/mcp", "server_label": "deepwiki"}], "Using the deepwiki tool, what does the xai-org/xai-sdk-python repo do? Max 10 words.", "test_tools mcp")87    mc = next(o for o in r["output"] if o["type"] == "mcp_call")88    assert mc["server_label"] == "deepwiki" and mc["output"]89    assert r["usage"]["server_side_tool_usage_details"]["mcp_calls"] >= 1909192@pytest.mark.run_expensive_tests93def test_file_search_tool_shape(xai, models):94    st, col, _ = xai("POST", "/v1/collections", {"collection_name": "atlas-test", "field_definitions": []}, note="test_tools collection create")95    assert st == 20096    cid = col["collection_id"]97    try:98        r = _run_tool(xai, models, [{"type": "file_search", "vector_store_ids": [cid]}], "What is in my documents? Three words.", "test_tools file_search", ["file_search_call.results"])99        fs = next(o for o in r["output"] if o["type"] == "file_search_call")100        assert fs["status"] in ("completed", "failed") and isinstance(fs["queries"], list)101    finally:102        xai("DELETE", f"/v1/collections/{cid}", note="test_tools collection delete")103