demoprep / tools /ts_mcp_answer_probe.py
mike boone
fix: detect invalid answer readiness before MCP liveboards
d3049d4
Raw
History Blame Contribute Delete
7.8 kB
#!/usr/bin/env python3
"""Probe ThoughtSpot direct answer API and MCP getAnswer for an existing model."""
from __future__ import annotations
import argparse
import asyncio
import json
import os
from typing import Any
import requests
from dotenv import load_dotenv
import yaml
DEFAULT_QUESTIONS = [
"sum [Ad Revenue Usd]",
"average [Arpu Usd]",
"average [Fill Rate Pct]",
"sum [Ad Revenue Usd] by [Month Start Date].monthly",
]
def _env(name: str, default: str = "") -> str:
return os.getenv(name, default).strip()
def _auth(base_url: str, username: str, secret_key: str) -> tuple[requests.Session, str]:
response = requests.post(
f"{base_url}/api/rest/2.0/auth/token/full",
json={
"username": username,
"secret_key": secret_key,
"validity_time_in_sec": 3600,
},
timeout=30,
)
response.raise_for_status()
token = response.json()["token"]
session = requests.Session()
session.headers.update(
{
"Authorization": f"Bearer {token}",
"Content-Type": "application/json",
"Accept": "application/json",
}
)
return session, token
def _shape(data: Any) -> dict:
if not isinstance(data, dict):
return {"type": type(data).__name__}
tokens = data.get("tokens")
display_tokens = data.get("display_tokens")
shaped = {
"keys": sorted(data.keys()),
"session_identifier": bool(data.get("session_identifier")),
"tokens_type": type(tokens).__name__,
"tokens_len": len(tokens) if isinstance(tokens, (str, list)) else None,
"display_tokens_type": type(display_tokens).__name__,
"display_tokens_len": len(display_tokens) if isinstance(display_tokens, (str, list)) else None,
"visualization_type": data.get("visualization_type"),
"generation_number": data.get("generation_number"),
"message_type": data.get("message_type"),
}
if "error" in data:
shaped["error"] = data.get("error")
return shaped
def export_model_columns(session: requests.Session, base_url: str, model_id: str) -> list[dict]:
response = session.post(
f"{base_url}/api/rest/2.0/metadata/tml/export",
json={"metadata": [{"identifier": model_id}], "export_associated": False},
timeout=60,
)
if response.status_code != 200:
print(f"Model TML export failed: HTTP {response.status_code} {response.text[:300]}")
return []
payload = response.json()
if not payload:
return []
edoc = payload[0].get("edoc") or ""
parsed = yaml.safe_load(edoc) or {}
columns = []
for col in (parsed.get("model") or {}).get("columns", []) or []:
props = col.get("properties") or {}
columns.append(
{
"name": col.get("name"),
"column_type": props.get("column_type"),
"aggregation": props.get("aggregation"),
"calendar": props.get("calendar"),
}
)
return columns
def search_model_details(session: requests.Session, base_url: str, model_id: str) -> dict:
response = session.post(
f"{base_url}/api/rest/2.0/metadata/search",
json={
"metadata": [{"type": "LOGICAL_TABLE", "identifier": model_id}],
"record_size": 1,
"include_details": True,
},
timeout=60,
)
if response.status_code != 200:
return {"http_status": response.status_code, "body": response.text[:500]}
data = response.json()
if not data:
return {}
return data[0]
def probe_direct(session: requests.Session, base_url: str, model_id: str, question: str) -> dict:
response = session.post(
f"{base_url}/api/rest/2.0/ai/answer/create",
json={"query": question, "metadata_identifier": model_id},
timeout=90,
)
result = {
"http_status": response.status_code,
"question": question,
}
try:
data = response.json()
except Exception:
result["body"] = response.text[:500]
return result
result["shape"] = _shape(data)
return result
async def probe_mcp(base_url: str, token: str, model_id: str, questions: list[str]) -> list[dict]:
from mcp import ClientSession
from mcp.client.streamable_http import streamablehttp_client
host = base_url.replace("https://", "").replace("http://", "").rstrip("/")
endpoint = "https://agent.thoughtspot.app/bearer/mcp"
headers = {"Authorization": f"Bearer {token}@{host}"}
results = []
async with streamablehttp_client(endpoint, headers=headers) as (read, write, _):
async with ClientSession(read, write) as mcp_session:
await mcp_session.initialize()
for question in questions:
record = {"question": question}
try:
answer_result = await mcp_session.call_tool(
"getAnswer",
{"question": question, "datasourceId": model_id},
)
text = answer_result.content[0].text if answer_result.content else ""
record["raw_text_prefix"] = text[:300]
try:
data = json.loads(text)
except Exception as exc:
record["parse_error"] = f"{type(exc).__name__}: {exc}"
else:
record["shape"] = _shape(data)
except Exception as exc:
record["error"] = f"{type(exc).__name__}: {exc}"
results.append(record)
return results
def main() -> int:
load_dotenv()
parser = argparse.ArgumentParser(description=__doc__)
parser.add_argument("model_id")
parser.add_argument("--ts-url", default=_env("TS_ENV_1_URL") or _env("THOUGHTSPOT_URL"))
parser.add_argument("--ts-username", default=_env("TEST_USER") or _env("THOUGHTSPOT_USERNAME"))
parser.add_argument("--ts-secret-key", default=_env("TS_ENV_1_KEY_VAR") or _env("THOUGHTSPOT_TRUSTED_AUTH_KEY"))
parser.add_argument("--question", action="append", dest="questions")
args = parser.parse_args()
if not args.ts_url or not args.ts_username or not args.ts_secret_key:
raise SystemExit("Missing TS auth. Provide --ts-url/--ts-username/--ts-secret-key or env vars.")
base_url = args.ts_url.rstrip("/")
questions = args.questions or DEFAULT_QUESTIONS
print(f"ThoughtSpot URL: {base_url}")
print(f"ThoughtSpot user: {args.ts_username}")
print(f"Model: {args.model_id}")
print(f"Questions: {len(questions)}")
session, token = _auth(base_url, args.ts_username, args.ts_secret_key)
columns = export_model_columns(session, base_url, args.model_id)
print("\nModel columns:")
if columns:
for col in columns:
marker = "MEASURE" if col.get("column_type") == "MEASURE" else ("DATE" if col.get("calendar") else "ATTRIBUTE")
print(f"- {col.get('name')} [{marker}]")
else:
detail = search_model_details(session, base_url, args.model_id)
header = detail.get("metadata_header") or {}
print(f"- detail keys: {sorted(detail.keys()) if detail else []}")
print(f"- header: name={header.get('name')} id={header.get('id')} type={header.get('type')}")
print("\nDirect answer API:")
for question in questions:
print(json.dumps(probe_direct(session, base_url, args.model_id, question), indent=2))
print("\nMCP getAnswer:")
for record in asyncio.run(probe_mcp(base_url, token, args.model_id, questions)):
print(json.dumps(record, indent=2))
return 0
if __name__ == "__main__":
raise SystemExit(main())