mirror of
https://github.com/turnstonelabs/turnstone.git
synced 2026-08-12 23:12:23 -06:00
Refactor detect_model into shared function, auto-detect context window
Move detect_model() to turnstone.core.model_registry as a single implementation replacing duplicates in cli.py, server.py, and eval.py. Auto-detects context window from backend metadata (meta.n_ctx_train) when available, falling back to the 131072 default otherwise. Also replace vLLM-specific references in help text and comments with generic "OpenAI-compatible API" / "model server" language.
This commit is contained in:
+18
-22
@@ -690,24 +690,11 @@ def _handle_cluster_command(cmd_line: str, console_url: str | None, auth_token:
|
||||
# ─── Model auto-detection ─────────────────────────────────────────────────
|
||||
|
||||
|
||||
def detect_model(client: OpenAI) -> str:
|
||||
"""Auto-detect the model from vLLM's /v1/models endpoint."""
|
||||
try:
|
||||
models = client.models.list()
|
||||
model_ids = [m.id for m in models.data]
|
||||
if not model_ids:
|
||||
print(red("No models found at server. Use --model to specify."))
|
||||
sys.exit(1)
|
||||
if len(model_ids) == 1:
|
||||
return model_ids[0]
|
||||
# Multiple models -- pick first, but inform user
|
||||
print(f"Available models: {', '.join(model_ids)}")
|
||||
print(f"Using: {bold(model_ids[0])} (override with --model)")
|
||||
return model_ids[0]
|
||||
except Exception as e:
|
||||
print(red(f"Could not connect to server: {e}"))
|
||||
print("Is vLLM running? Start it or use --base-url to point elsewhere.")
|
||||
sys.exit(1)
|
||||
def detect_model(client: OpenAI) -> tuple[str, int | None]:
|
||||
"""Auto-detect model — delegates to :func:`turnstone.core.model_registry.detect_model`."""
|
||||
from turnstone.core.model_registry import detect_model as _detect
|
||||
|
||||
return _detect(client)
|
||||
|
||||
|
||||
# ─── Main ──────────────────────────────────────────────────────────────────
|
||||
@@ -715,7 +702,7 @@ def detect_model(client: OpenAI) -> str:
|
||||
|
||||
def main() -> None:
|
||||
parser = argparse.ArgumentParser(
|
||||
description="Interactive CLI for vLLM models with tool calling.",
|
||||
description="Interactive CLI for OpenAI-compatible models with tool calling.",
|
||||
formatter_class=argparse.RawDescriptionHelpFormatter,
|
||||
epilog=textwrap.dedent("""\
|
||||
Examples:
|
||||
@@ -727,7 +714,7 @@ def main() -> None:
|
||||
parser.add_argument(
|
||||
"--base-url",
|
||||
default="http://localhost:8000/v1",
|
||||
help="vLLM API base URL (default: http://localhost:8000/v1)",
|
||||
help="OpenAI-compatible API base URL (default: http://localhost:8000/v1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
@@ -851,7 +838,16 @@ def main() -> None:
|
||||
base_url=args.base_url,
|
||||
api_key=api_key,
|
||||
)
|
||||
model = args.model or detect_model(client)
|
||||
if args.model:
|
||||
model = args.model
|
||||
detected_ctx = None
|
||||
else:
|
||||
model, detected_ctx = detect_model(client)
|
||||
|
||||
# Use detected context window when the user hasn't explicitly set one
|
||||
context_window = args.context_window
|
||||
if detected_ctx and context_window == 131072: # default unchanged
|
||||
context_window = detected_ctx
|
||||
|
||||
# Build model registry (reads [models.*] sections from config.toml)
|
||||
from turnstone.core.model_registry import load_model_registry
|
||||
@@ -860,7 +856,7 @@ def main() -> None:
|
||||
base_url=args.base_url,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
context_window=args.context_window,
|
||||
context_window=context_window,
|
||||
)
|
||||
|
||||
# Initialize MCP client (connects to configured MCP servers, if any)
|
||||
|
||||
@@ -210,3 +210,42 @@ def load_model_registry(
|
||||
fallback=fallback,
|
||||
agent_model=agent_model,
|
||||
)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
# Model auto-detection
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def detect_model(client: OpenAI, log_fn: Any = print) -> tuple[str, int | None]:
|
||||
"""Auto-detect the model and context window from the /v1/models endpoint.
|
||||
|
||||
Returns ``(model_id, context_window)`` where *context_window* is
|
||||
``None`` when the backend does not expose ``meta.n_ctx_train``.
|
||||
|
||||
Calls ``log_fn`` for informational messages (defaults to ``print``).
|
||||
Raises ``SystemExit`` on failure.
|
||||
"""
|
||||
try:
|
||||
models = client.models.list()
|
||||
if not models.data:
|
||||
log_fn("Error: No models found at server. Use --model to specify.")
|
||||
raise SystemExit(1)
|
||||
m = models.data[0]
|
||||
if len(models.data) > 1:
|
||||
log_fn(f"Available models: {', '.join(x.id for x in models.data)}")
|
||||
log_fn(f"Using: {m.id} (override with --model)")
|
||||
# Extract context window from backend metadata (llama.cpp, vLLM, etc.)
|
||||
ctx: int | None = None
|
||||
meta = m.model_dump().get("meta")
|
||||
if isinstance(meta, dict):
|
||||
n_ctx = meta.get("n_ctx_train")
|
||||
if isinstance(n_ctx, int) and n_ctx > 0:
|
||||
ctx = n_ctx
|
||||
return m.id, ctx
|
||||
except SystemExit:
|
||||
raise
|
||||
except Exception as e:
|
||||
log_fn(f"Error: Could not connect to server: {e}")
|
||||
log_fn("Is the model server running? Start it or use --base-url to point elsewhere.")
|
||||
raise SystemExit(1) from e
|
||||
|
||||
@@ -595,7 +595,7 @@ class ChatSession:
|
||||
"""Stream response, dispatching tokens to the UI as they arrive.
|
||||
|
||||
Handles two reasoning delivery mechanisms:
|
||||
1. vLLM's `reasoning_content` field (when --reasoning-parser is set)
|
||||
1. The `reasoning_content` field (e.g. vLLM with --reasoning-parser)
|
||||
2. <think>...</think> tags in regular content (common default)
|
||||
|
||||
Calls self.ui.on_thinking_stop() on the first received delta.
|
||||
@@ -715,7 +715,7 @@ class ChatSession:
|
||||
if parts:
|
||||
self.ui.on_info(f"{GRAY}[delta: {', '.join(parts)}]{RESET}")
|
||||
|
||||
# Path 1: reasoning field (vLLM sends as "reasoning" or "reasoning_content")
|
||||
# Path 1: reasoning field (sent as "reasoning" or "reasoning_content")
|
||||
rc = getattr(delta, "reasoning", None) or getattr(delta, "reasoning_content", None)
|
||||
if rc:
|
||||
_stop_spinner_once()
|
||||
|
||||
+3
-13
@@ -964,7 +964,9 @@ def run_optimization(
|
||||
)
|
||||
|
||||
if not model:
|
||||
model = _detect_model(client)
|
||||
from turnstone.core.model_registry import detect_model
|
||||
|
||||
model, _ = detect_model(client)
|
||||
|
||||
# Load test cases
|
||||
with open(test_file) as f:
|
||||
@@ -1127,18 +1129,6 @@ def run_optimization(
|
||||
return results
|
||||
|
||||
|
||||
def _detect_model(client: OpenAI) -> str:
|
||||
"""Auto-detect the model from the API."""
|
||||
try:
|
||||
models = client.models.list()
|
||||
model_ids = [m.id for m in models.data]
|
||||
if model_ids:
|
||||
return model_ids[0]
|
||||
except Exception:
|
||||
pass
|
||||
raise SystemExit("Could not auto-detect model. Use --model to specify.")
|
||||
|
||||
|
||||
# ─── CLI ─────────────────────────────────────────────────────────────────────
|
||||
|
||||
|
||||
|
||||
+18
-20
@@ -922,23 +922,11 @@ class ThreadedHTTPServer(ThreadingMixIn, HTTPServer):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def detect_model(client: OpenAI) -> str:
|
||||
"""Auto-detect the model from vLLM's /v1/models endpoint."""
|
||||
try:
|
||||
models = client.models.list()
|
||||
model_ids = [m.id for m in models.data]
|
||||
if not model_ids:
|
||||
print("Error: No models found at server. Use --model to specify.")
|
||||
sys.exit(1)
|
||||
if len(model_ids) == 1:
|
||||
return model_ids[0]
|
||||
print(f"Available models: {', '.join(model_ids)}")
|
||||
print(f"Using: {model_ids[0]} (override with --model)")
|
||||
return model_ids[0]
|
||||
except Exception as e:
|
||||
print(f"Error: Could not connect to server: {e}")
|
||||
print("Is vLLM running? Start it or use --base-url to point elsewhere.")
|
||||
sys.exit(1)
|
||||
def detect_model(client: OpenAI) -> tuple[str, int | None]:
|
||||
"""Auto-detect model — delegates to :func:`turnstone.core.model_registry.detect_model`."""
|
||||
from turnstone.core.model_registry import detect_model as _detect
|
||||
|
||||
return _detect(client)
|
||||
|
||||
|
||||
# ---------------------------------------------------------------------------
|
||||
@@ -1002,7 +990,7 @@ def main() -> None:
|
||||
parser.add_argument(
|
||||
"--base-url",
|
||||
default="http://localhost:8000/v1",
|
||||
help="vLLM API base URL (default: http://localhost:8000/v1)",
|
||||
help="OpenAI-compatible API base URL (default: http://localhost:8000/v1)",
|
||||
)
|
||||
parser.add_argument(
|
||||
"--model",
|
||||
@@ -1182,7 +1170,17 @@ def main() -> None:
|
||||
base_url=args.base_url,
|
||||
api_key=api_key,
|
||||
)
|
||||
model = args.model or detect_model(client)
|
||||
if args.model:
|
||||
model = args.model
|
||||
detected_ctx = None
|
||||
else:
|
||||
model, detected_ctx = detect_model(client)
|
||||
|
||||
# Use detected context window when the user hasn't explicitly set one
|
||||
context_window = args.context_window
|
||||
if detected_ctx and context_window == 131072: # default unchanged
|
||||
context_window = detected_ctx
|
||||
print(f"Context window: {context_window:,} (detected from backend)")
|
||||
|
||||
# Build model registry (reads [models.*] sections from config.toml)
|
||||
from turnstone.core.model_registry import load_model_registry
|
||||
@@ -1191,7 +1189,7 @@ def main() -> None:
|
||||
base_url=args.base_url,
|
||||
api_key=api_key,
|
||||
model=model,
|
||||
context_window=args.context_window,
|
||||
context_window=context_window,
|
||||
)
|
||||
|
||||
# Initialize MCP client (connects to configured MCP servers, if any)
|
||||
|
||||
Reference in New Issue
Block a user