mirror of
https://github.com/pewdiepie-archdaemon/odysseus.git
synced 2026-08-05 02:45:28 +00:00
Merge 16f7e058b4 into 20e7fc0164
This commit is contained in:
commit
7cc7020961
10 changed files with 706 additions and 17 deletions
|
|
@ -477,6 +477,7 @@ class ProviderAuthSession(TimestampMixin, Base):
|
|||
base_url = Column(String, nullable=False)
|
||||
access_token = Column(EncryptedText, nullable=True)
|
||||
refresh_token = Column(EncryptedText, nullable=True)
|
||||
chatgpt_account_id = Column(EncryptedText, nullable=True)
|
||||
last_refresh = Column(DateTime, nullable=True)
|
||||
auth_mode = Column(String, nullable=True)
|
||||
|
||||
|
|
@ -974,6 +975,30 @@ def _migrate_add_provider_auth_id_column():
|
|||
pass
|
||||
|
||||
|
||||
def _migrate_add_provider_auth_session_account_id_column():
|
||||
"""Add ChatGPT account metadata to provider_auth_sessions if missing."""
|
||||
import sqlite3
|
||||
db_path = DATABASE_URL.replace("sqlite:///", "")
|
||||
if not os.path.exists(db_path):
|
||||
return
|
||||
conn = None
|
||||
try:
|
||||
conn = sqlite3.connect(db_path)
|
||||
cursor = conn.execute("PRAGMA table_info(provider_auth_sessions)")
|
||||
columns = [row[1] for row in cursor.fetchall()]
|
||||
if columns and "chatgpt_account_id" not in columns:
|
||||
conn.execute("ALTER TABLE provider_auth_sessions ADD COLUMN chatgpt_account_id TEXT")
|
||||
conn.commit()
|
||||
logging.getLogger(__name__).info("Migrated: added 'chatgpt_account_id' column to provider_auth_sessions")
|
||||
except Exception as e:
|
||||
logging.getLogger(__name__).warning(f"provider_auth_sessions.chatgpt_account_id migration failed: {e}")
|
||||
finally:
|
||||
try:
|
||||
conn.close()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _migrate_add_model_type_column():
|
||||
"""Add model_type column to model_endpoints if it doesn't exist."""
|
||||
import sqlite3
|
||||
|
|
@ -1934,6 +1959,7 @@ def init_db():
|
|||
_migrate_add_model_endpoint_refresh_columns()
|
||||
_migrate_add_model_endpoint_owner_column()
|
||||
_migrate_add_provider_auth_id_column()
|
||||
_migrate_add_provider_auth_session_account_id_column()
|
||||
_migrate_add_supports_tools_column()
|
||||
_migrate_add_task_run_model_column()
|
||||
_migrate_add_owner_column()
|
||||
|
|
|
|||
|
|
@ -55,6 +55,8 @@ def _provision_endpoint(tokens: Dict, owner: Optional[str]) -> Dict:
|
|||
auth.base_url = base
|
||||
auth.access_token = access_token
|
||||
auth.refresh_token = refresh_token
|
||||
auth.chatgpt_account_id = chatgpt_subscription.chatgpt_account_id_from_tokens(tokens)
|
||||
chatgpt_subscription.remember_chatgpt_account_id(access_token, auth.chatgpt_account_id)
|
||||
auth.last_refresh = utcnow_naive()
|
||||
auth.auth_mode = "chatgpt"
|
||||
|
||||
|
|
@ -82,7 +84,7 @@ def _provision_endpoint(tokens: Dict, owner: Optional[str]) -> Dict:
|
|||
ep.api_key = None
|
||||
ep.provider_auth_id = auth.id
|
||||
ep.is_enabled = True
|
||||
ep.supports_tools = False
|
||||
ep.supports_tools = True
|
||||
ep.model_type = "llm"
|
||||
ep.endpoint_kind = "api"
|
||||
ep.model_refresh_mode = "manual"
|
||||
|
|
|
|||
|
|
@ -3625,6 +3625,12 @@ async def stream_agent_loop(
|
|||
_db.close()
|
||||
except Exception as _e:
|
||||
logger.debug(f"endpoint supports_tools lookup failed: {_e}")
|
||||
try:
|
||||
from src.chatgpt_subscription import is_chatgpt_subscription_base as _is_chatgpt_subscription_base
|
||||
if _is_chatgpt_subscription_base(endpoint_url or ""):
|
||||
_endpoint_supports = True
|
||||
except Exception:
|
||||
pass
|
||||
_model_supports_tools = any(kw in _model_lc for kw in (
|
||||
"gpt-4", "gpt-5", "gpt-o", "claude", "gemini", "gemma",
|
||||
"qwen3", "qwen2.5", "mixtral", "mistral", "llama-3.1", "llama-3.2",
|
||||
|
|
|
|||
|
|
@ -29,6 +29,8 @@ CHATGPT_OAUTH_REDIRECT_URI = f"{CHATGPT_OAUTH_ISSUER}/deviceauth/callback"
|
|||
CHATGPT_ACCESS_TOKEN_REFRESH_SKEW_SECONDS = 120
|
||||
_AUTH_REFRESH_LOCKS: dict[str, threading.Lock] = {}
|
||||
_AUTH_REFRESH_LOCKS_GUARD = threading.Lock()
|
||||
_ACCOUNT_IDS_BY_ACCESS_TOKEN: dict[str, str] = {}
|
||||
_ACCOUNT_IDS_GUARD = threading.Lock()
|
||||
|
||||
|
||||
def _database_handles():
|
||||
|
|
@ -75,7 +77,25 @@ def is_chatgpt_subscription_base(url: str) -> bool:
|
|||
)
|
||||
|
||||
|
||||
def chatgpt_headers(access_token: Optional[str]) -> Dict[str, str]:
|
||||
def _remember_account_for_access_token(access_token: Optional[str], account_id: Optional[str]) -> None:
|
||||
if not access_token or not account_id:
|
||||
return
|
||||
with _ACCOUNT_IDS_GUARD:
|
||||
_ACCOUNT_IDS_BY_ACCESS_TOKEN[str(access_token)] = str(account_id)
|
||||
|
||||
|
||||
def remember_chatgpt_account_id(access_token: Optional[str], account_id: Optional[str]) -> None:
|
||||
_remember_account_for_access_token(access_token, account_id)
|
||||
|
||||
|
||||
def _remembered_account_for_access_token(access_token: Optional[str]) -> Optional[str]:
|
||||
if not access_token:
|
||||
return None
|
||||
with _ACCOUNT_IDS_GUARD:
|
||||
return _ACCOUNT_IDS_BY_ACCESS_TOKEN.get(str(access_token))
|
||||
|
||||
|
||||
def chatgpt_headers(access_token: Optional[str], account_id: Optional[str] = None) -> Dict[str, str]:
|
||||
headers = {
|
||||
"Accept": "application/json, text/event-stream",
|
||||
"Origin": "https://chatgpt.com",
|
||||
|
|
@ -84,6 +104,9 @@ def chatgpt_headers(access_token: Optional[str]) -> Dict[str, str]:
|
|||
}
|
||||
if access_token:
|
||||
headers["Authorization"] = f"Bearer {access_token}"
|
||||
resolved_account_id = account_id or _remembered_account_for_access_token(access_token)
|
||||
if resolved_account_id:
|
||||
headers["ChatGPT-Account-Id"] = str(resolved_account_id)
|
||||
return headers
|
||||
|
||||
|
||||
|
|
@ -243,6 +266,33 @@ def _decode_jwt_payload(token: str) -> Dict[str, Any]:
|
|||
return payload if isinstance(payload, dict) else {}
|
||||
|
||||
|
||||
def _chatgpt_account_id_from_jwt(token: Optional[str]) -> Optional[str]:
|
||||
if not token:
|
||||
return None
|
||||
try:
|
||||
payload = _decode_jwt_payload(token)
|
||||
except Exception:
|
||||
return None
|
||||
auth_claims = payload.get("https://api.openai.com/auth")
|
||||
if isinstance(auth_claims, dict):
|
||||
account_id = auth_claims.get("chatgpt_account_id")
|
||||
if isinstance(account_id, str) and account_id.strip():
|
||||
return account_id.strip()
|
||||
for key in ("chatgpt_account_id", "account_id"):
|
||||
account_id = payload.get(key)
|
||||
if isinstance(account_id, str) and account_id.strip():
|
||||
return account_id.strip()
|
||||
return None
|
||||
|
||||
|
||||
def chatgpt_account_id_from_tokens(tokens: Dict[str, Any]) -> Optional[str]:
|
||||
for key in ("account_id", "chatgpt_account_id"):
|
||||
account_id = tokens.get(key)
|
||||
if isinstance(account_id, str) and account_id.strip():
|
||||
return account_id.strip()
|
||||
return _chatgpt_account_id_from_jwt(tokens.get("id_token")) or _chatgpt_account_id_from_jwt(tokens.get("access_token"))
|
||||
|
||||
|
||||
def access_token_is_expiring(access_token: str, skew_seconds: int = CHATGPT_ACCESS_TOKEN_REFRESH_SKEW_SECONDS) -> bool:
|
||||
try:
|
||||
exp = int(_decode_jwt_payload(access_token).get("exp") or 0)
|
||||
|
|
@ -276,16 +326,23 @@ def resolve_runtime_credentials(auth_id: str, owner: Optional[str] = None, *, fo
|
|||
row.access_token = refreshed["access_token"]
|
||||
if refreshed.get("refresh_token"):
|
||||
row.refresh_token = refreshed["refresh_token"]
|
||||
account_id = chatgpt_account_id_from_tokens(refreshed)
|
||||
if account_id:
|
||||
row.chatgpt_account_id = account_id
|
||||
row.last_refresh = utcnow_naive()
|
||||
db.commit()
|
||||
db.refresh(row)
|
||||
access_token = row.access_token or ""
|
||||
|
||||
account_id = (getattr(row, "chatgpt_account_id", None) or "").strip() or chatgpt_account_id_from_tokens({"access_token": access_token})
|
||||
if account_id:
|
||||
_remember_account_for_access_token(access_token, account_id)
|
||||
return {
|
||||
"provider": CHATGPT_SUBSCRIPTION_PROVIDER,
|
||||
"base_url": (row.base_url or DEFAULT_CHATGPT_SUBSCRIPTION_BASE_URL).rstrip("/"),
|
||||
"api_key": access_token,
|
||||
"auth_mode": row.auth_mode or "chatgpt",
|
||||
"chatgpt_account_id": account_id,
|
||||
}
|
||||
finally:
|
||||
db.close()
|
||||
|
|
@ -300,16 +357,84 @@ def to_http_exception(exc: Exception) -> HTTPException:
|
|||
|
||||
|
||||
def build_responses_input(messages: list[dict]) -> list[dict]:
|
||||
def _content_text(content: Any) -> str:
|
||||
if isinstance(content, list):
|
||||
return "\n".join(
|
||||
str(part.get("text") or part.get("content") or "")
|
||||
for part in content
|
||||
if isinstance(part, dict)
|
||||
)
|
||||
return "" if content is None else str(content)
|
||||
|
||||
def _tool_call_item(tool_call: dict) -> Optional[dict]:
|
||||
if not isinstance(tool_call, dict):
|
||||
return None
|
||||
fn = tool_call.get("function") if isinstance(tool_call.get("function"), dict) else {}
|
||||
name = str(fn.get("name") or tool_call.get("name") or "").strip()
|
||||
if not name:
|
||||
return None
|
||||
arguments = fn.get("arguments", tool_call.get("arguments", "{}"))
|
||||
if not isinstance(arguments, str):
|
||||
arguments = json.dumps(arguments if arguments is not None else {})
|
||||
call_id = str(tool_call.get("id") or tool_call.get("call_id") or "").strip()
|
||||
if not call_id:
|
||||
call_id = f"call_{len(input_items)}"
|
||||
return {
|
||||
"type": "function_call",
|
||||
"call_id": call_id,
|
||||
"name": name,
|
||||
"arguments": arguments,
|
||||
}
|
||||
|
||||
input_items: list[dict] = []
|
||||
for msg in messages or []:
|
||||
role = msg.get("role") or "user"
|
||||
if role == "tool":
|
||||
role = "user"
|
||||
content = msg.get("content")
|
||||
if isinstance(content, list):
|
||||
text = "\n".join(str(part.get("text") or part.get("content") or "") for part in content if isinstance(part, dict))
|
||||
else:
|
||||
text = "" if content is None else str(content)
|
||||
call_id = str(msg.get("tool_call_id") or msg.get("call_id") or "").strip()
|
||||
if call_id:
|
||||
input_items.append({
|
||||
"type": "function_call_output",
|
||||
"call_id": call_id,
|
||||
"output": _content_text(msg.get("content")),
|
||||
})
|
||||
continue
|
||||
text = _content_text(msg.get("content"))
|
||||
input_type = "output_text" if role == "assistant" else "input_text"
|
||||
input_items.append({"role": role, "content": [{"type": input_type, "text": text}]})
|
||||
if text or role != "assistant" or not msg.get("tool_calls"):
|
||||
input_items.append({"role": role, "content": [{"type": input_type, "text": text}]})
|
||||
if role == "assistant":
|
||||
for tool_call in msg.get("tool_calls") or []:
|
||||
item = _tool_call_item(tool_call)
|
||||
if item:
|
||||
input_items.append(item)
|
||||
return input_items
|
||||
|
||||
|
||||
def build_responses_tools(tools: list[dict] | None) -> list[dict]:
|
||||
"""Convert OpenAI chat-completions tool schemas to Responses function tools."""
|
||||
response_tools: list[dict] = []
|
||||
for tool in tools or []:
|
||||
if not isinstance(tool, dict) or tool.get("type") != "function":
|
||||
continue
|
||||
fn = tool.get("function")
|
||||
if isinstance(fn, dict):
|
||||
name = str(fn.get("name") or "").strip()
|
||||
if not name:
|
||||
continue
|
||||
item: dict[str, Any] = {
|
||||
"type": "function",
|
||||
"name": name,
|
||||
"parameters": fn.get("parameters") or {},
|
||||
}
|
||||
if fn.get("description"):
|
||||
item["description"] = str(fn["description"])
|
||||
if "strict" in fn:
|
||||
item["strict"] = bool(fn["strict"])
|
||||
elif "strict" in tool:
|
||||
item["strict"] = bool(tool["strict"])
|
||||
response_tools.append(item)
|
||||
continue
|
||||
name = str(tool.get("name") or "").strip()
|
||||
if name:
|
||||
response_tools.append(dict(tool))
|
||||
return response_tools
|
||||
|
|
|
|||
106
src/llm_core.py
106
src/llm_core.py
|
|
@ -1110,8 +1110,9 @@ def _build_chatgpt_responses_payload(
|
|||
max_tokens: int,
|
||||
*,
|
||||
stream: bool = False,
|
||||
tools: Optional[List[Dict]] = None,
|
||||
) -> Dict:
|
||||
from src.chatgpt_subscription import build_responses_input
|
||||
from src.chatgpt_subscription import build_responses_input, build_responses_tools
|
||||
|
||||
conversation = [msg for msg in (messages or []) if (msg.get("role") or "") != "system"]
|
||||
payload: Dict = {
|
||||
|
|
@ -1123,6 +1124,9 @@ def _build_chatgpt_responses_payload(
|
|||
}
|
||||
if not _restricts_temperature(model):
|
||||
payload["temperature"] = temperature
|
||||
response_tools = build_responses_tools(tools)
|
||||
if response_tools:
|
||||
payload["tools"] = response_tools
|
||||
# ChatGPT Subscription Codex API does not support max_output_tokens —
|
||||
# passing it returns HTTP 400 "Unsupported parameter: max_output_tokens".
|
||||
# Do not include it in the payload.
|
||||
|
|
@ -2209,7 +2213,10 @@ async def _stream_llm_inner(url: str, model: str, messages: List[Dict], temperat
|
|||
elif provider == "chatgpt-subscription":
|
||||
target_url = _normalize_chatgpt_subscription_url(url)
|
||||
h = _provider_headers(provider, headers)
|
||||
payload = _build_chatgpt_responses_payload(model, messages_copy, temperature, max_tokens, stream=True)
|
||||
payload = _build_chatgpt_responses_payload(
|
||||
model, messages_copy, temperature, max_tokens,
|
||||
stream=True, tools=tools,
|
||||
)
|
||||
else:
|
||||
target_url = _normalize_openai_chat_url(url)
|
||||
payload = {
|
||||
|
|
@ -2266,6 +2273,65 @@ async def _stream_llm_inner(url: str, model: str, messages: List[Dict], temperat
|
|||
event_name = ""
|
||||
input_tokens = 0
|
||||
output_tokens = 0
|
||||
_responses_tool_calls: Dict[str, Dict] = {}
|
||||
_responses_tool_order: List[str] = []
|
||||
_responses_tool_aliases: Dict[str, str] = {}
|
||||
|
||||
def _alias_values(data: Dict, item: Optional[Dict] = None) -> List[str]:
|
||||
aliases: List[str] = []
|
||||
src = item if isinstance(item, dict) else {}
|
||||
for value in (src.get("id"), src.get("call_id"), data.get("item_id"), data.get("output_item_id")):
|
||||
if value is not None:
|
||||
aliases.append(str(value))
|
||||
if data.get("output_index") is not None:
|
||||
aliases.append(f"output_index:{data.get('output_index')}")
|
||||
return aliases
|
||||
|
||||
def _tool_call_key(data: Dict, item: Optional[Dict] = None) -> str:
|
||||
for alias in _alias_values(data, item):
|
||||
if alias in _responses_tool_aliases:
|
||||
return _responses_tool_aliases[alias]
|
||||
aliases = _alias_values(data, item)
|
||||
key = aliases[0] if aliases else f"response_tool_{len(_responses_tool_order)}"
|
||||
for alias in aliases:
|
||||
_responses_tool_aliases[alias] = key
|
||||
return key
|
||||
|
||||
def _remember_tool_item(data: Dict, item: Dict) -> Optional[Dict]:
|
||||
if not isinstance(item, dict):
|
||||
return None
|
||||
if item.get("type") not in {"function_call", "custom_tool_call"}:
|
||||
return None
|
||||
key = _tool_call_key(data, item)
|
||||
if key not in _responses_tool_calls:
|
||||
_responses_tool_calls[key] = {"id": "", "name": "", "arguments": ""}
|
||||
_responses_tool_order.append(key)
|
||||
for alias in _alias_values(data, item):
|
||||
_responses_tool_aliases[alias] = key
|
||||
call = _responses_tool_calls[key]
|
||||
call_id = item.get("call_id") or item.get("id")
|
||||
if call_id:
|
||||
call["id"] = str(call_id)
|
||||
if item.get("name"):
|
||||
call["name"] = str(item["name"])
|
||||
if "arguments" in item and item.get("arguments") is not None:
|
||||
call["arguments"] = str(item.get("arguments") or "")
|
||||
return call
|
||||
|
||||
def _emit_responses_tool_calls() -> Optional[str]:
|
||||
calls = [
|
||||
_responses_tool_calls[key]
|
||||
for key in _responses_tool_order
|
||||
if _responses_tool_calls.get(key, {}).get("name")
|
||||
]
|
||||
if not calls:
|
||||
return None
|
||||
for idx, call in enumerate(calls):
|
||||
if not call.get("id"):
|
||||
call["id"] = f"call_{idx}"
|
||||
call.setdefault("arguments", "")
|
||||
return f'data: {json.dumps({"type": "tool_calls", "calls": calls})}\n\n'
|
||||
|
||||
try:
|
||||
client = _get_http_client()
|
||||
async with client.stream('POST', target_url, json=payload, headers=h, timeout=stream_timeout) as r:
|
||||
|
|
@ -2286,6 +2352,12 @@ async def _stream_llm_inner(url: str, model: str, messages: List[Dict], temperat
|
|||
raw = line[5:].strip()
|
||||
if not raw:
|
||||
continue
|
||||
if raw == "[DONE]":
|
||||
tc_event = _emit_responses_tool_calls()
|
||||
if tc_event:
|
||||
yield tc_event
|
||||
yield "data: [DONE]\n\n"
|
||||
return
|
||||
try:
|
||||
data = json.loads(raw)
|
||||
except json.JSONDecodeError:
|
||||
|
|
@ -2299,8 +2371,33 @@ async def _stream_llm_inner(url: str, model: str, messages: List[Dict], temperat
|
|||
yield _degenerate
|
||||
return
|
||||
yield f'data: {json.dumps({"delta": delta})}\n\n'
|
||||
elif evt in {"response.output_item.added", "response.output_item.done"}:
|
||||
_remember_tool_item(data, data.get("item") or data.get("output_item") or {})
|
||||
elif evt in {"response.function_call_arguments.delta", "response.custom_tool_call_input.delta"}:
|
||||
key = _tool_call_key(data)
|
||||
if key not in _responses_tool_calls:
|
||||
_responses_tool_calls[key] = {"id": "", "name": "", "arguments": ""}
|
||||
_responses_tool_order.append(key)
|
||||
delta = data.get("delta") or ""
|
||||
_responses_tool_calls[key]["arguments"] += str(delta)
|
||||
name = _responses_tool_calls[key].get("name")
|
||||
if delta and name in ("create_document", "update_document", "edit_document"):
|
||||
yield f'data: {json.dumps({"type": "tool_call_delta", "index": _responses_tool_order.index(key), "name": name, "arg_delta": delta})}\n\n'
|
||||
elif evt in {"response.function_call_arguments.done", "response.custom_tool_call_input.done"}:
|
||||
key = _tool_call_key(data)
|
||||
if key not in _responses_tool_calls:
|
||||
_responses_tool_calls[key] = {"id": "", "name": "", "arguments": ""}
|
||||
_responses_tool_order.append(key)
|
||||
if data.get("arguments") is not None:
|
||||
_responses_tool_calls[key]["arguments"] = str(data.get("arguments") or "")
|
||||
elif evt == "response.completed":
|
||||
usage = (data.get("response") or {}).get("usage") or data.get("usage") or {}
|
||||
response = data.get("response") or {}
|
||||
for item in response.get("output") or []:
|
||||
_remember_tool_item(data, item)
|
||||
tc_event = _emit_responses_tool_calls()
|
||||
if tc_event:
|
||||
yield tc_event
|
||||
usage = response.get("usage") or data.get("usage") or {}
|
||||
input_tokens = usage.get("input_tokens") or usage.get("prompt_tokens") or input_tokens
|
||||
output_tokens = usage.get("output_tokens") or usage.get("completion_tokens") or output_tokens
|
||||
if input_tokens or output_tokens:
|
||||
|
|
@ -2312,6 +2409,9 @@ async def _stream_llm_inner(url: str, model: str, messages: List[Dict], temperat
|
|||
text = err.get("message") if isinstance(err, dict) else str(err or "ChatGPT Subscription request failed")
|
||||
yield f'event: error\ndata: {json.dumps({"status": 502, "text": text})}\n\n'
|
||||
return
|
||||
tc_event = _emit_responses_tool_calls()
|
||||
if tc_event:
|
||||
yield tc_event
|
||||
yield "data: [DONE]\n\n"
|
||||
except (httpx.ConnectError, httpx.ConnectTimeout) as e:
|
||||
_cooled = _mark_host_dead(target_url)
|
||||
|
|
|
|||
|
|
@ -15,9 +15,10 @@ import logging
|
|||
import os
|
||||
import pathlib
|
||||
import re
|
||||
import shlex
|
||||
import sys
|
||||
import time
|
||||
from typing import Any, Awaitable, Callable, Dict, Optional, Tuple
|
||||
from typing import Any, Awaitable, Callable, Dict, List, Optional, Tuple
|
||||
|
||||
|
||||
|
||||
|
|
@ -501,6 +502,34 @@ def _promote_image_fields(result: Dict) -> None:
|
|||
|
||||
|
||||
_BG_MARKERS = {"#!bg", "#bg", "# bg", "#background", "# background", "@background", "# @background"}
|
||||
_WORKSPACE_SHELL_MUTATION_CMD_FALLBACK_RE = re.compile(
|
||||
r"(^|[;&|]\s*)(cp|copy|copy-item|mv|move|rename|ren|touch|tee)\b|"
|
||||
r"(^|[;&|]\s*)(sed\s+-i|perl\s+-pi|awk\s+-i)\b",
|
||||
re.IGNORECASE | re.MULTILINE,
|
||||
)
|
||||
_WORKSPACE_SHELL_MUTATION_COMMANDS = {
|
||||
"cp",
|
||||
"copy",
|
||||
"copy-item",
|
||||
"mv",
|
||||
"move",
|
||||
"rename",
|
||||
"ren",
|
||||
"touch",
|
||||
"tee",
|
||||
}
|
||||
_WORKSPACE_SHELL_IN_PLACE_COMMANDS = {"awk", "perl", "sed"}
|
||||
_WORKSPACE_SHELL_COMMAND_POSITION_WORDS = {"then", "do", "else", "elif"}
|
||||
_WORKSPACE_SHELL_ASSIGNMENT_RE = re.compile(r"^[A-Za-z_][A-Za-z0-9_]*=.*$")
|
||||
_WORKSPACE_SHELL_PUNCTUATION = set("(){};<>|&")
|
||||
_WORKSPACE_SHELL_HEREDOC_TOKENS = {"<<", "<<-"}
|
||||
_WORKSPACE_SHELL_OUTPUT_REDIRECT_TOKENS = {">", ">>", ">|", "&>", "&>>", ">&", ">>&"}
|
||||
_WORKSPACE_SHELL_OUTPUT_REDIRECT_RE = re.compile(
|
||||
r"^(?P<fd>\d*)?(?P<op>&>>|>>&|>>|>\||&>|>&|>)(?P<target>.*)$"
|
||||
)
|
||||
_WORKSPACE_SHELL_FD_TARGET_RE = re.compile(r"^&?\d+$|^&?-$")
|
||||
_WORKSPACE_SHELL_DEV_FD_RE = re.compile(r"^/(?:dev/fd|proc/self/fd)/\d+$")
|
||||
_WORKSPACE_SHELL_SINK_TARGETS = {"/dev/null", "/dev/stdout", "/dev/stderr", os.devnull.lower()}
|
||||
|
||||
|
||||
def _split_bg_marker(content: str):
|
||||
|
|
@ -516,6 +545,182 @@ def _split_bg_marker(content: str):
|
|||
return False, content
|
||||
|
||||
|
||||
def _workspace_shell_tokens(command: str) -> Optional[List[str]]:
|
||||
lexer = shlex.shlex(command, posix=True, punctuation_chars=True)
|
||||
lexer.whitespace_split = True
|
||||
lexer.commenters = ""
|
||||
try:
|
||||
return list(lexer)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
|
||||
def _workspace_shell_is_punctuation_token(token: str) -> bool:
|
||||
return bool(token) and all(char in _WORKSPACE_SHELL_PUNCTUATION for char in token)
|
||||
|
||||
|
||||
def _workspace_shell_starts_command_position(token: str) -> bool:
|
||||
if not _workspace_shell_is_punctuation_token(token):
|
||||
return token.lower() in _WORKSPACE_SHELL_COMMAND_POSITION_WORDS
|
||||
if ">" in token or "<" in token:
|
||||
return False
|
||||
return any(char in token for char in ";&|({")
|
||||
|
||||
|
||||
def _workspace_shell_command_name(token: str) -> str:
|
||||
command = token.strip("`").replace("\\", "/").rstrip("/")
|
||||
return command.rsplit("/", 1)[-1].lower()
|
||||
|
||||
|
||||
def _workspace_shell_perl_in_place_option(option: str) -> bool:
|
||||
if not option.startswith("-") or option == "--":
|
||||
return False
|
||||
flags = option[1:]
|
||||
return (
|
||||
option.startswith("-i")
|
||||
or option.startswith("-pi")
|
||||
or ("p" in flags and "i" in flags)
|
||||
)
|
||||
|
||||
|
||||
def _workspace_shell_in_place_option(command: str, option: str) -> bool:
|
||||
option = option.lower()
|
||||
if command == "perl":
|
||||
return _workspace_shell_perl_in_place_option(option)
|
||||
if command in {"awk", "sed"}:
|
||||
return (
|
||||
option == "--in-place"
|
||||
or option.startswith("--in-place=")
|
||||
or option.startswith("-i")
|
||||
)
|
||||
return False
|
||||
|
||||
|
||||
def _workspace_shell_command_is_mutating(tokens: List[str], index: int) -> bool:
|
||||
command = _workspace_shell_command_name(tokens[index])
|
||||
if command in _WORKSPACE_SHELL_MUTATION_COMMANDS:
|
||||
return True
|
||||
if command not in _WORKSPACE_SHELL_IN_PLACE_COMMANDS:
|
||||
return False
|
||||
|
||||
for token in tokens[index + 1 :]:
|
||||
if _workspace_shell_starts_command_position(token):
|
||||
break
|
||||
if token == "--":
|
||||
break
|
||||
if _workspace_shell_in_place_option(command, token):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def _workspace_shell_has_mutation_command(command: str) -> bool:
|
||||
tokens = _workspace_shell_tokens(command)
|
||||
if tokens is None:
|
||||
return bool(_WORKSPACE_SHELL_MUTATION_CMD_FALLBACK_RE.search(command))
|
||||
|
||||
expect_command = True
|
||||
for index, token in enumerate(tokens):
|
||||
if _workspace_shell_starts_command_position(token):
|
||||
expect_command = True
|
||||
continue
|
||||
if not expect_command:
|
||||
continue
|
||||
if _WORKSPACE_SHELL_ASSIGNMENT_RE.match(token):
|
||||
continue
|
||||
if _workspace_shell_command_is_mutating(tokens, index):
|
||||
return True
|
||||
expect_command = False
|
||||
return False
|
||||
|
||||
|
||||
def _workspace_shell_redirect_target_is_safe(target: str, workspace: str) -> bool:
|
||||
target = (target or "").strip()
|
||||
if not target:
|
||||
return False
|
||||
if _WORKSPACE_SHELL_FD_TARGET_RE.match(target):
|
||||
return True
|
||||
|
||||
lowered = target.lower()
|
||||
if lowered in _WORKSPACE_SHELL_SINK_TARGETS:
|
||||
return True
|
||||
if _WORKSPACE_SHELL_DEV_FD_RE.match(target):
|
||||
return True
|
||||
|
||||
expanded = os.path.expandvars(os.path.expanduser(target))
|
||||
if not os.path.isabs(expanded):
|
||||
return False
|
||||
|
||||
try:
|
||||
workspace_real = os.path.realpath(workspace)
|
||||
target_real = os.path.realpath(expanded)
|
||||
return os.path.commonpath([workspace_real, target_real]) != workspace_real
|
||||
except (OSError, ValueError):
|
||||
return False
|
||||
|
||||
|
||||
def _workspace_shell_redirects_to_workspace(command: str, workspace: str) -> bool:
|
||||
tokens = _workspace_shell_tokens(command)
|
||||
if tokens is None:
|
||||
# If tokenization cannot determine quoting, keep the old safety posture
|
||||
# for literal redirects while avoiding quoted/comparison false positives
|
||||
# in valid shell.
|
||||
return bool(re.search(r"(^|[\s;&|])(?:\d*)>{1,2}\S*", command) or "<<" in command)
|
||||
|
||||
i = 0
|
||||
while i < len(tokens):
|
||||
token = tokens[i]
|
||||
if token in _WORKSPACE_SHELL_HEREDOC_TOKENS or token.startswith("<<"):
|
||||
return True
|
||||
|
||||
if token.isdigit() and i + 1 < len(tokens):
|
||||
next_token = tokens[i + 1]
|
||||
if next_token in _WORKSPACE_SHELL_OUTPUT_REDIRECT_TOKENS:
|
||||
target = tokens[i + 2] if i + 2 < len(tokens) else ""
|
||||
if not _workspace_shell_redirect_target_is_safe(target, workspace):
|
||||
return True
|
||||
i += 3
|
||||
continue
|
||||
|
||||
if token in _WORKSPACE_SHELL_OUTPUT_REDIRECT_TOKENS:
|
||||
target = tokens[i + 1] if i + 1 < len(tokens) else ""
|
||||
if not _workspace_shell_redirect_target_is_safe(target, workspace):
|
||||
return True
|
||||
i += 2
|
||||
continue
|
||||
|
||||
match = _WORKSPACE_SHELL_OUTPUT_REDIRECT_RE.match(token)
|
||||
if match:
|
||||
target = match.group("target")
|
||||
if match.group("op") == ">&" and not target:
|
||||
target = tokens[i + 1] if i + 1 < len(tokens) else ""
|
||||
i += 1
|
||||
if not _workspace_shell_redirect_target_is_safe(target, workspace):
|
||||
return True
|
||||
i += 1
|
||||
|
||||
return False
|
||||
|
||||
|
||||
def _workspace_shell_write_block_reason(tool: str, content: str) -> Optional[str]:
|
||||
workspace = get_active_workspace()
|
||||
if tool != "bash" or not workspace:
|
||||
return None
|
||||
_, command = _split_bg_marker(content or "")
|
||||
if not command.strip():
|
||||
return None
|
||||
if not (
|
||||
_workspace_shell_has_mutation_command(command)
|
||||
or _workspace_shell_redirects_to_workspace(command, workspace)
|
||||
):
|
||||
return None
|
||||
return (
|
||||
"Workspace file changes must use `write_file` for creates/full rewrites "
|
||||
"or `edit_file` for targeted edits. Shell is still available for read-only "
|
||||
"diagnostics, but redirection/heredocs/tee/cp/mv/touch/in-place edits are "
|
||||
"blocked while a workspace is active."
|
||||
)
|
||||
|
||||
|
||||
async def _direct_fallback(
|
||||
tool: str,
|
||||
content: str,
|
||||
|
|
@ -711,6 +916,12 @@ async def _execute_tool_block_impl(
|
|||
logger.warning("Public tool policy blocked owner=%r tool=%s", owner, tool)
|
||||
return desc, result
|
||||
|
||||
workspace_shell_block = _workspace_shell_write_block_reason(tool, content)
|
||||
if workspace_shell_block:
|
||||
desc = f"{tool}: BLOCKED"
|
||||
result = {"error": workspace_shell_block, "exit_code": 1}
|
||||
logger.info("Workspace shell write blocked for tool=%s", tool)
|
||||
return desc, result
|
||||
|
||||
# Background execution: a `bash` block whose first line is the `#!bg`
|
||||
# marker runs DETACHED — returns a job id immediately so the chat stream
|
||||
|
|
|
|||
|
|
@ -1,5 +1,6 @@
|
|||
"""DB-backed ChatGPT Subscription endpoint provisioning tests."""
|
||||
|
||||
import base64
|
||||
import json
|
||||
|
||||
import pytest
|
||||
|
|
@ -10,6 +11,23 @@ from core.database import Base, ModelEndpoint, ProviderAuthSession
|
|||
import routes.chatgpt_subscription_routes as csr
|
||||
|
||||
|
||||
def _jwt(payload):
|
||||
def enc(obj):
|
||||
raw = json.dumps(obj, separators=(",", ":")).encode("utf-8")
|
||||
return base64.urlsafe_b64encode(raw).decode("ascii").rstrip("=")
|
||||
|
||||
return f"{enc({'alg': 'none'})}.{enc(payload)}.sig"
|
||||
|
||||
|
||||
def _id_token(account_id="acct_test"):
|
||||
return _jwt({
|
||||
"https://api.openai.com/auth": {
|
||||
"chatgpt_account_id": account_id,
|
||||
"chatgpt_user_id": "user_test",
|
||||
}
|
||||
})
|
||||
|
||||
|
||||
def _mem_db(monkeypatch):
|
||||
engine = create_engine("sqlite:///:memory:")
|
||||
Base.metadata.create_all(bind=engine)
|
||||
|
|
@ -25,7 +43,11 @@ def test_provision_creates_owner_scoped_auth_session_and_endpoint(monkeypatch):
|
|||
TestSessionLocal = _mem_db(monkeypatch)
|
||||
monkeypatch.setattr(csr.chatgpt_subscription, "fetch_available_models", lambda token: ["gpt-5.5", "o4-mini"])
|
||||
|
||||
res = csr._provision_endpoint({"access_token": "AT", "refresh_token": "RT"}, "alice")
|
||||
res = csr._provision_endpoint({
|
||||
"access_token": "AT",
|
||||
"refresh_token": "RT",
|
||||
"id_token": _id_token("acct_alice"),
|
||||
}, "alice")
|
||||
|
||||
assert res["name"] == "ChatGPT Subscription"
|
||||
assert res["base_url"] == csr.chatgpt_subscription.DEFAULT_CHATGPT_SUBSCRIPTION_BASE_URL
|
||||
|
|
@ -40,14 +62,16 @@ def test_provision_creates_owner_scoped_auth_session_and_endpoint(monkeypatch):
|
|||
assert auth.provider == csr.chatgpt_subscription.CHATGPT_SUBSCRIPTION_PROVIDER
|
||||
assert auth.access_token == "AT"
|
||||
assert auth.refresh_token == "RT"
|
||||
assert auth.chatgpt_account_id == "acct_alice"
|
||||
assert auth.auth_mode == "chatgpt"
|
||||
assert csr.chatgpt_subscription.chatgpt_headers("AT")["ChatGPT-Account-Id"] == "acct_alice"
|
||||
assert ep is not None
|
||||
assert ep.owner == "alice"
|
||||
assert ep.api_key is None
|
||||
assert ep.provider_auth_id == auth.id
|
||||
assert ep.endpoint_kind == "api"
|
||||
assert ep.model_refresh_mode == "manual"
|
||||
assert ep.supports_tools is False
|
||||
assert ep.supports_tools is True
|
||||
assert json.loads(ep.cached_models) == ["gpt-5.5", "o4-mini"]
|
||||
finally:
|
||||
db.close()
|
||||
|
|
|
|||
|
|
@ -45,7 +45,7 @@ class _FakeClient:
|
|||
return _FakeStreamCtx(self._lines)
|
||||
|
||||
|
||||
def _drive(monkeypatch, lines, model="gemini-3.1-pro-preview-customtools"):
|
||||
def _drive(monkeypatch, lines, model="gemini-3.1-pro-preview-customtools", url=None):
|
||||
"""Run stream_llm against a canned SSE line list; return parsed events."""
|
||||
monkeypatch.setattr(llm_core, "_get_http_client", lambda: _FakeClient(lines))
|
||||
monkeypatch.setattr(llm_core, "_is_host_dead", lambda u: False)
|
||||
|
|
@ -55,7 +55,7 @@ def _drive(monkeypatch, lines, model="gemini-3.1-pro-preview-customtools"):
|
|||
async def run():
|
||||
events = []
|
||||
async for chunk in llm_core.stream_llm(
|
||||
"https://generativelanguage.googleapis.com/v1beta/openai/chat/completions",
|
||||
url or "https://generativelanguage.googleapis.com/v1beta/openai/chat/completions",
|
||||
model,
|
||||
[{"role": "user", "content": "hi"}],
|
||||
headers={"Authorization": "Bearer k"},
|
||||
|
|
@ -77,6 +77,12 @@ def _sse(delta):
|
|||
return "data: " + json.dumps({"choices": [{"delta": delta}]})
|
||||
|
||||
|
||||
def _event_data(event_type, payload):
|
||||
body = dict(payload)
|
||||
body.setdefault("type", event_type)
|
||||
return [f"event: {event_type}", "data: " + json.dumps(body)]
|
||||
|
||||
|
||||
def test_parallel_calls_with_null_index_do_not_collide(monkeypatch):
|
||||
# Two parallel calls, each complete in one delta, both with index=None
|
||||
# (exactly what Gemini's OpenAI-compat layer emits). Only the first carries
|
||||
|
|
@ -106,6 +112,49 @@ def test_parallel_calls_with_null_index_do_not_collide(monkeypatch):
|
|||
assert "extra_content" not in by_name["bash"]
|
||||
|
||||
|
||||
def test_chatgpt_subscription_responses_function_call_stream(monkeypatch):
|
||||
lines = []
|
||||
lines += _event_data("response.output_item.added", {
|
||||
"output_index": 0,
|
||||
"item": {
|
||||
"type": "function_call",
|
||||
"id": "fc_read",
|
||||
"call_id": "call_read",
|
||||
"name": "read_file",
|
||||
"arguments": "",
|
||||
},
|
||||
})
|
||||
lines += _event_data("response.function_call_arguments.delta", {
|
||||
"output_index": 0,
|
||||
"item_id": "fc_read",
|
||||
"delta": '{"path": "',
|
||||
})
|
||||
lines += _event_data("response.function_call_arguments.delta", {
|
||||
"output_index": 0,
|
||||
"item_id": "fc_read",
|
||||
"delta": '/workspace/README.txt"}',
|
||||
})
|
||||
lines += _event_data("response.completed", {
|
||||
"response": {"usage": {"input_tokens": 7, "output_tokens": 2}},
|
||||
})
|
||||
|
||||
events = _drive(
|
||||
monkeypatch,
|
||||
lines,
|
||||
model="gpt-5.3-codex-spark",
|
||||
url="https://chatgpt.com/backend-api/codex",
|
||||
)
|
||||
|
||||
calls = next(e["calls"] for e in events if e.get("type") == "tool_calls")
|
||||
assert calls == [{
|
||||
"id": "call_read",
|
||||
"name": "read_file",
|
||||
"arguments": '{"path": "/workspace/README.txt"}',
|
||||
}]
|
||||
usage = next(e["data"] for e in events if e.get("type") == "usage")
|
||||
assert usage == {"input_tokens": 7, "output_tokens": 2}
|
||||
|
||||
|
||||
def test_single_call_chunked_arguments_still_accumulate(monkeypatch):
|
||||
# Conformant OpenAI style: index present, arguments streamed in pieces.
|
||||
lines = [
|
||||
|
|
|
|||
|
|
@ -9,6 +9,7 @@ import httpx
|
|||
import pytest
|
||||
|
||||
from src import llm_core
|
||||
from src import chatgpt_subscription
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
|
|
@ -131,3 +132,71 @@ def test_chatgpt_subscription_payload_omits_max_output_tokens_when_zero():
|
|||
)
|
||||
|
||||
assert "max_output_tokens" not in payload
|
||||
|
||||
|
||||
def test_chatgpt_subscription_payload_converts_function_tools():
|
||||
payload = llm_core._build_chatgpt_responses_payload(
|
||||
"gpt-5.3-codex-spark",
|
||||
[{"role": "user", "content": "Read the README"}],
|
||||
temperature=0.2,
|
||||
max_tokens=0,
|
||||
tools=[{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "read_file",
|
||||
"description": "Read a file",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"path": {"type": "string"}},
|
||||
"required": ["path"],
|
||||
},
|
||||
},
|
||||
}],
|
||||
)
|
||||
|
||||
assert payload["tools"] == [{
|
||||
"type": "function",
|
||||
"name": "read_file",
|
||||
"description": "Read a file",
|
||||
"parameters": {
|
||||
"type": "object",
|
||||
"properties": {"path": {"type": "string"}},
|
||||
"required": ["path"],
|
||||
},
|
||||
}]
|
||||
|
||||
|
||||
def test_chatgpt_subscription_input_preserves_native_tool_turns():
|
||||
items = chatgpt_subscription.build_responses_input([
|
||||
{
|
||||
"role": "assistant",
|
||||
"content": None,
|
||||
"tool_calls": [{
|
||||
"id": "call_read",
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": "read_file",
|
||||
"arguments": '{"path": "/workspace/README.txt"}',
|
||||
},
|
||||
}],
|
||||
},
|
||||
{
|
||||
"role": "tool",
|
||||
"tool_call_id": "call_read",
|
||||
"content": "README contents",
|
||||
},
|
||||
])
|
||||
|
||||
assert items == [
|
||||
{
|
||||
"type": "function_call",
|
||||
"call_id": "call_read",
|
||||
"name": "read_file",
|
||||
"arguments": '{"path": "/workspace/README.txt"}',
|
||||
},
|
||||
{
|
||||
"type": "function_call_output",
|
||||
"call_id": "call_read",
|
||||
"output": "README contents",
|
||||
},
|
||||
]
|
||||
|
|
|
|||
|
|
@ -20,6 +20,7 @@ import pytest
|
|||
from src.tool_execution import (
|
||||
_AGENT_WORKDIR,
|
||||
_active_workspace,
|
||||
_workspace_shell_write_block_reason,
|
||||
_resolve_search_root,
|
||||
_resolve_tool_path,
|
||||
_resolve_tool_path_in_workspace,
|
||||
|
|
@ -260,6 +261,82 @@ async def test_glob_skips_sensitive_files_in_workspace(ws, admin):
|
|||
assert r["exit_code"] == 0 and "No files" in r["output"]
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", [
|
||||
"awk '$3 > 100 {print $1}' data.csv",
|
||||
"cat data.json | jq '.items[] | select(.size > 5)'",
|
||||
'echo "use > to redirect"',
|
||||
"ls -la > /dev/null 2>&1",
|
||||
"grep -rn 'a -> b' src/",
|
||||
'python -c "print(1 > 0)"',
|
||||
"git log --oneline | head -20",
|
||||
"diff <(sort a.txt) <(sort b.txt)",
|
||||
])
|
||||
def test_workspace_shell_guard_allows_read_only_redirect_syntax(ws, command):
|
||||
token = _active_workspace.set(ws)
|
||||
try:
|
||||
assert _workspace_shell_write_block_reason("bash", command) is None
|
||||
finally:
|
||||
_active_workspace.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", [
|
||||
"grep -E 'mv|cp' log.txt",
|
||||
"awk '/mv|cp/' file",
|
||||
'echo "a;cp b"',
|
||||
])
|
||||
def test_workspace_shell_guard_allows_quoted_mutation_words(ws, command):
|
||||
token = _active_workspace.set(ws)
|
||||
try:
|
||||
assert _workspace_shell_write_block_reason("bash", command) is None
|
||||
finally:
|
||||
_active_workspace.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", [
|
||||
"cp secret.txt out.txt",
|
||||
"touch note.txt",
|
||||
"tee out.txt",
|
||||
"echo ok && cp a b",
|
||||
"(cp secret.txt out.txt)",
|
||||
"$(mv a.txt b.txt)",
|
||||
"`cp a b`",
|
||||
"{ cp a b; }",
|
||||
"sed -i 's/a/b/' file",
|
||||
"perl -pi -e 's/a/b/' file",
|
||||
"awk -i inplace '{print}' file",
|
||||
])
|
||||
def test_workspace_shell_guard_blocks_tokenized_mutation_commands(ws, command):
|
||||
token = _active_workspace.set(ws)
|
||||
try:
|
||||
assert _workspace_shell_write_block_reason("bash", command)
|
||||
finally:
|
||||
_active_workspace.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.parametrize("command", [
|
||||
"printf 'x' > note.txt",
|
||||
"printf 'x' >> note.txt",
|
||||
"printf 'x' 1> note.txt",
|
||||
"printf 'x' 2> error.log",
|
||||
"printf 'x' &> out.log",
|
||||
])
|
||||
def test_workspace_shell_guard_blocks_workspace_redirect_targets(ws, command):
|
||||
token = _active_workspace.set(ws)
|
||||
try:
|
||||
assert _workspace_shell_write_block_reason("bash", command)
|
||||
finally:
|
||||
_active_workspace.reset(token)
|
||||
|
||||
|
||||
def test_workspace_shell_guard_blocks_absolute_workspace_redirect_target(ws):
|
||||
target = os.path.join(ws, "absolute-note.txt")
|
||||
token = _active_workspace.set(ws)
|
||||
try:
|
||||
assert _workspace_shell_write_block_reason("bash", f"printf 'x' > {target}")
|
||||
finally:
|
||||
_active_workspace.reset(token)
|
||||
|
||||
|
||||
@pytest.mark.asyncio
|
||||
async def test_subprocess_cwd_is_workspace_e2e(ws, admin):
|
||||
"""python tool runs with cwd = workspace (OS-agnostic probe)."""
|
||||
|
|
|
|||
Loading…
Add table
Reference in a new issue