mirror of
https://github.com/openswarm-ai/openswarm.git
synced 2026-09-30 13:34:50 +02:00
[haik]: refactor: systematize naming conventions across 20+ backend modules. Private module-level vars renamed _foo to P_FOO, private funcs _foo to p_foo, cross-module public funcs drop leading underscores entirely (e.g. _discover_mcp_tools_http -> discover_mcp_tools_http, _load_all -> load_all_tools, _classify_services -> classify_services). Local vars inside functions lose underscores (_backend -> backend_dir, _port -> port). Unnecessary import aliases removed (import json as _json -> import json). Imports reorganized to proper source modules (derive_mcp_config from mcp_config, oauth refresh funcs from oauth_tokens). Added Public comments documenting cross-module API boundaries. Extracted PENDING_OAUTH to oauth_state.py to break circular import. No behavioral changes.
This commit is contained in:
@@ -14,17 +14,18 @@ from backend.apps.agents.core.models import (
|
||||
)
|
||||
from backend.apps.agents.core.ws_manager import ws_manager
|
||||
from backend.apps.settings.store import load_settings
|
||||
from backend.apps.tools_lib.oauth_tokens import (
|
||||
refresh_google_token,
|
||||
refresh_airtable_token,
|
||||
refresh_hubspot_token,
|
||||
)
|
||||
from backend.apps.tools_lib.tools_lib import (
|
||||
_load_all as load_all_tools,
|
||||
derive_mcp_config,
|
||||
load_all_tools,
|
||||
load_builtin_permissions,
|
||||
load_trusted_sensitive_paths,
|
||||
refresh_airtable_token,
|
||||
refresh_google_token,
|
||||
refresh_hubspot_token,
|
||||
save_trusted_sensitive_paths,
|
||||
)
|
||||
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
|
||||
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name, derive_mcp_config
|
||||
from backend.apps.agents.core.error_classify import (
|
||||
is_auth_error,
|
||||
is_free_trial_exhausted,
|
||||
|
||||
@@ -373,8 +373,8 @@ async def subscriptions_connect(body: dict):
|
||||
result = await start_oauth(provider)
|
||||
|
||||
if result.get("flow") == "authorization_code" and result.get("state"):
|
||||
from backend.main import _pending_oauth
|
||||
_pending_oauth[result["state"]] = {
|
||||
from backend.apps.oauth_state import PENDING_OAUTH
|
||||
PENDING_OAUTH[result["state"]] = {
|
||||
"provider": provider,
|
||||
"code_verifier": result.get("code_verifier", ""),
|
||||
"redirect_uri": result.get("redirect_uri", ""),
|
||||
|
||||
@@ -11,7 +11,7 @@ from typing import Any
|
||||
from backend.apps.agents.providers.registry import resolve_aux_model
|
||||
from backend.apps.settings.credentials import get_anthropic_client_for_model
|
||||
from backend.apps.settings.store import load_settings
|
||||
from backend.apps.tools_lib.tools_lib import _load_all as load_all_tools
|
||||
from backend.apps.tools_lib.tools_lib import load_all_tools
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
@@ -2,7 +2,7 @@ from typing import Callable
|
||||
|
||||
from backend.apps.modes.modes import load_mode
|
||||
from backend.apps.tools_lib.tools_lib import (
|
||||
_load_all as load_all_tools,
|
||||
load_all_tools,
|
||||
)
|
||||
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
|
||||
from backend.apps.agents.manager.prompt.tool_catalog import get_denied_tool_names, is_fully_denied
|
||||
|
||||
@@ -19,10 +19,10 @@ from backend.apps.settings.models import AppSettings, DEFAULT_SYSTEM_PROMPT
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
SETTINGS_FILE = os.path.join(DATA_DIR, "settings.json")
|
||||
P_SETTINGS_FILE = os.path.join(DATA_DIR, "settings.json")
|
||||
|
||||
|
||||
def _migrate_legacy_fields(raw: dict) -> dict:
|
||||
def p_migrate_legacy_fields(raw: dict) -> dict:
|
||||
"""Translate deprecated pre-launch field names ('managed', 'openswarm_auth_token') into production schema."""
|
||||
if raw.get("connection_mode") == "managed":
|
||||
raw["connection_mode"] = "openswarm-pro"
|
||||
@@ -31,7 +31,7 @@ def _migrate_legacy_fields(raw: dict) -> dict:
|
||||
return raw
|
||||
|
||||
|
||||
def _coerce_settings(raw: dict) -> AppSettings:
|
||||
def p_coerce_settings(raw: dict) -> AppSettings:
|
||||
"""Build AppSettings, surviving a settings.json written by a different app
|
||||
version. Unknown fields are already ignored by pydantic; the case this guards
|
||||
is a field whose TYPE drifted across versions (e.g. a list that is now a
|
||||
@@ -55,12 +55,12 @@ def _coerce_settings(raw: dict) -> AppSettings:
|
||||
return AppSettings()
|
||||
|
||||
|
||||
def _preserve_corrupt_settings() -> None:
|
||||
def p_preserve_corrupt_settings() -> None:
|
||||
"""Move an unparseable settings.json aside so boot proceeds on defaults while
|
||||
the original stays recoverable (the next save would otherwise overwrite it)."""
|
||||
try:
|
||||
backup = SETTINGS_FILE + ".corrupt"
|
||||
os.replace(SETTINGS_FILE, backup)
|
||||
backup = P_SETTINGS_FILE + ".corrupt"
|
||||
os.replace(P_SETTINGS_FILE, backup)
|
||||
logger.warning("settings.json was unparseable; preserved at %s", backup)
|
||||
except OSError:
|
||||
pass
|
||||
@@ -70,13 +70,13 @@ def _preserve_corrupt_settings() -> None:
|
||||
# so even a hand-edited file or an unexpected writer is picked up immediately. A stat
|
||||
# skips the open+parse+validate that Defender turns into 5-50ms on Windows. Copies on
|
||||
# both sides keep handler isolation: callers mutate their copy, never the cache.
|
||||
_cached_settings: AppSettings | None = None
|
||||
_cached_sig: tuple[int, int] | None = None
|
||||
P_CACHED_SETTINGS: AppSettings | None = None
|
||||
P_CACHED_SIG: tuple[int, int] | None = None
|
||||
|
||||
|
||||
def _settings_sig() -> tuple[int, int] | None:
|
||||
def p_settings_sig() -> tuple[int, int] | None:
|
||||
try:
|
||||
st = os.stat(SETTINGS_FILE)
|
||||
st = os.stat(P_SETTINGS_FILE)
|
||||
return (st.st_mtime_ns, st.st_size)
|
||||
except OSError:
|
||||
return None
|
||||
@@ -86,38 +86,38 @@ def load_settings() -> AppSettings:
|
||||
"""Load settings from JSON file, returning defaults if not found. Never raises
|
||||
on a corrupt or version-mismatched file: a single bad settings.json must not
|
||||
brick boot (it is read at startup, by the settings endpoint, and per dispatch)."""
|
||||
global _cached_settings, _cached_sig
|
||||
sig = _settings_sig()
|
||||
if sig is not None and _cached_settings is not None and sig == _cached_sig:
|
||||
return _cached_settings.model_copy(deep=True)
|
||||
if os.path.exists(SETTINGS_FILE):
|
||||
global P_CACHED_SETTINGS, P_CACHED_SIG
|
||||
sig = p_settings_sig()
|
||||
if sig is not None and P_CACHED_SETTINGS is not None and sig == P_CACHED_SIG:
|
||||
return P_CACHED_SETTINGS.model_copy(deep=True)
|
||||
if os.path.exists(P_SETTINGS_FILE):
|
||||
try:
|
||||
with open(SETTINGS_FILE) as f:
|
||||
with open(P_SETTINGS_FILE) as f:
|
||||
raw = json.load(f)
|
||||
except (json.JSONDecodeError, OSError, ValueError):
|
||||
_preserve_corrupt_settings()
|
||||
p_preserve_corrupt_settings()
|
||||
return AppSettings()
|
||||
if not isinstance(raw, dict):
|
||||
# Valid JSON but not an object (e.g. a bare list/number); unusable.
|
||||
_preserve_corrupt_settings()
|
||||
p_preserve_corrupt_settings()
|
||||
return AppSettings()
|
||||
settings = _coerce_settings(_migrate_legacy_fields(raw))
|
||||
settings = p_coerce_settings(p_migrate_legacy_fields(raw))
|
||||
if settings.default_system_prompt is None:
|
||||
settings.default_system_prompt = DEFAULT_SYSTEM_PROMPT
|
||||
_cached_settings = settings.model_copy(deep=True)
|
||||
_cached_sig = sig
|
||||
P_CACHED_SETTINGS = settings.model_copy(deep=True)
|
||||
P_CACHED_SIG = sig
|
||||
return settings
|
||||
return AppSettings()
|
||||
|
||||
|
||||
# threading.Lock guards every SETTINGS_FILE write; works for sync paths and async run_in_executor paths.
|
||||
_settings_write_lock = threading.Lock()
|
||||
P_SETTINGS_WRITE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def atomic_write_settings(payload: dict) -> None:
|
||||
"""Atomic SETTINGS_FILE write; call via save_settings*, not directly."""
|
||||
global _cached_settings, _cached_sig
|
||||
with _settings_write_lock:
|
||||
global P_CACHED_SETTINGS, P_CACHED_SIG
|
||||
with P_SETTINGS_WRITE_LOCK:
|
||||
os.makedirs(DATA_DIR, exist_ok=True)
|
||||
fd, tmp = tempfile.mkstemp(prefix=".settings.", suffix=".tmp", dir=DATA_DIR)
|
||||
try:
|
||||
@@ -126,12 +126,12 @@ def atomic_write_settings(payload: dict) -> None:
|
||||
# Windows: Defender can briefly lock the destination; one retry handles every real case.
|
||||
for attempt in range(2):
|
||||
try:
|
||||
os.replace(tmp, SETTINGS_FILE)
|
||||
os.replace(tmp, P_SETTINGS_FILE)
|
||||
# Refresh the cache inside the lock so cache order matches disk order.
|
||||
_cached_settings = _coerce_settings(_migrate_legacy_fields(dict(payload)))
|
||||
if _cached_settings.default_system_prompt is None:
|
||||
_cached_settings.default_system_prompt = DEFAULT_SYSTEM_PROMPT
|
||||
_cached_sig = _settings_sig()
|
||||
P_CACHED_SETTINGS = p_coerce_settings(p_migrate_legacy_fields(dict(payload)))
|
||||
if P_CACHED_SETTINGS.default_system_prompt is None:
|
||||
P_CACHED_SETTINGS.default_system_prompt = DEFAULT_SYSTEM_PROMPT
|
||||
P_CACHED_SIG = p_settings_sig()
|
||||
return
|
||||
except PermissionError:
|
||||
if attempt == 1:
|
||||
|
||||
@@ -11,19 +11,19 @@ from backend.config.Apps import SubApp
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
REPO = "anthropics/skills"
|
||||
BRANCH = "main"
|
||||
RAW_BASE = f"https://raw.githubusercontent.com/{REPO}/{BRANCH}"
|
||||
MANIFEST_URL = f"{RAW_BASE}/.claude-plugin/marketplace.json"
|
||||
REFRESH_INTERVAL_S = 3600
|
||||
CONCURRENT_FETCHES = 15
|
||||
P_REPO = "anthropics/skills"
|
||||
P_BRANCH = "main"
|
||||
P_RAW_BASE = f"https://raw.githubusercontent.com/{P_REPO}/{P_BRANCH}"
|
||||
P_MANIFEST_URL = f"{P_RAW_BASE}/.claude-plugin/marketplace.json"
|
||||
P_REFRESH_INTERVAL_S = 3600
|
||||
P_CONCURRENT_FETCHES = 15
|
||||
|
||||
_cache: dict[str, dict] = {}
|
||||
_cache_updated_at: float = 0
|
||||
_refresh_task: Optional[asyncio.Task] = None
|
||||
P_CACHE: dict[str, dict] = {}
|
||||
P_CACHE_UPDATED_AT: float = 0
|
||||
P_REFRESH_TASK: Optional[asyncio.Task] = None
|
||||
|
||||
|
||||
def _parse_frontmatter(raw: str) -> tuple[dict, str]:
|
||||
def p_parse_frontmatter(raw: str) -> tuple[dict, str]:
|
||||
"""Split YAML frontmatter from markdown body."""
|
||||
if not raw.startswith("---"):
|
||||
return {}, raw
|
||||
@@ -40,12 +40,12 @@ def _parse_frontmatter(raw: str) -> tuple[dict, str]:
|
||||
return meta, body
|
||||
|
||||
|
||||
async def _fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]:
|
||||
async def p_fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]:
|
||||
"""Fetch the marketplace.json manifest and return (skill_folder, plugin_name) pairs.
|
||||
|
||||
Uses raw.githubusercontent.com; no GitHub API needed, no rate limiting.
|
||||
"""
|
||||
resp = await client.get(MANIFEST_URL)
|
||||
resp = await client.get(P_MANIFEST_URL)
|
||||
resp.raise_for_status()
|
||||
manifest = resp.json()
|
||||
|
||||
@@ -58,7 +58,7 @@ async def _fetch_skill_paths(client: httpx.AsyncClient) -> list[tuple[str, str]]
|
||||
return paths
|
||||
|
||||
|
||||
async def _fetch_one_skill(
|
||||
async def p_fetch_one_skill(
|
||||
client: httpx.AsyncClient,
|
||||
sem: asyncio.Semaphore,
|
||||
folder: str,
|
||||
@@ -66,7 +66,7 @@ async def _fetch_one_skill(
|
||||
) -> Optional[dict]:
|
||||
async with sem:
|
||||
try:
|
||||
resp = await client.get(f"{RAW_BASE}/{folder}/SKILL.md")
|
||||
resp = await client.get(f"{P_RAW_BASE}/{folder}/SKILL.md")
|
||||
if resp.status_code != 200:
|
||||
return None
|
||||
raw = resp.text
|
||||
@@ -74,7 +74,7 @@ async def _fetch_one_skill(
|
||||
logger.debug(f"Failed to fetch {folder}/SKILL.md: {exc}")
|
||||
return None
|
||||
|
||||
meta, body = _parse_frontmatter(raw)
|
||||
meta, body = p_parse_frontmatter(raw)
|
||||
name = meta.get("name", "")
|
||||
if not name:
|
||||
folder_name = folder.rsplit("/", 1)[-1]
|
||||
@@ -86,23 +86,23 @@ async def _fetch_one_skill(
|
||||
"content": body,
|
||||
"folder": folder,
|
||||
"category": plugin_name.replace("-", " ").replace("_", " ").title(),
|
||||
"repositoryUrl": f"https://github.com/{REPO}/tree/{BRANCH}/{folder}",
|
||||
"repositoryUrl": f"https://github.com/{P_REPO}/tree/{P_BRANCH}/{folder}",
|
||||
}
|
||||
|
||||
|
||||
async def _fetch_all_skills() -> dict[str, dict]:
|
||||
async def p_fetch_all_skills() -> dict[str, dict]:
|
||||
skills: dict[str, dict] = {}
|
||||
async with httpx.AsyncClient(timeout=30.0) as client:
|
||||
try:
|
||||
paths = await _fetch_skill_paths(client)
|
||||
paths = await p_fetch_skill_paths(client)
|
||||
except Exception as e:
|
||||
logger.warning(f"Skill registry manifest fetch failed: {e}")
|
||||
return skills
|
||||
|
||||
logger.info(f"Skill registry: found {len(paths)} skills in manifest, fetching content...")
|
||||
sem = asyncio.Semaphore(CONCURRENT_FETCHES)
|
||||
sem = asyncio.Semaphore(P_CONCURRENT_FETCHES)
|
||||
results = await asyncio.gather(
|
||||
*[_fetch_one_skill(client, sem, folder, plugin) for folder, plugin in paths]
|
||||
*[p_fetch_one_skill(client, sem, folder, plugin) for folder, plugin in paths]
|
||||
)
|
||||
for rec in results:
|
||||
if rec:
|
||||
@@ -112,26 +112,26 @@ async def _fetch_all_skills() -> dict[str, dict]:
|
||||
return skills
|
||||
|
||||
|
||||
async def _refresh_loop():
|
||||
global _cache, _cache_updated_at
|
||||
async def p_refresh_loop():
|
||||
global P_CACHE, P_CACHE_UPDATED_AT
|
||||
while True:
|
||||
try:
|
||||
_cache = await _fetch_all_skills()
|
||||
_cache_updated_at = time.time()
|
||||
P_CACHE = await p_fetch_all_skills()
|
||||
P_CACHE_UPDATED_AT = time.time()
|
||||
except Exception as e:
|
||||
logger.exception(f"Skill registry refresh error: {e}")
|
||||
await asyncio.sleep(REFRESH_INTERVAL_S)
|
||||
await asyncio.sleep(P_REFRESH_INTERVAL_S)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
async def skill_registry_lifespan():
|
||||
global _refresh_task
|
||||
_refresh_task = asyncio.create_task(_refresh_loop())
|
||||
global P_REFRESH_TASK
|
||||
P_REFRESH_TASK = asyncio.create_task(p_refresh_loop())
|
||||
yield
|
||||
if _refresh_task:
|
||||
_refresh_task.cancel()
|
||||
if P_REFRESH_TASK:
|
||||
P_REFRESH_TASK.cancel()
|
||||
try:
|
||||
await _refresh_task
|
||||
await P_REFRESH_TASK
|
||||
except asyncio.CancelledError:
|
||||
pass
|
||||
|
||||
@@ -142,13 +142,13 @@ skill_registry = SubApp("skill-registry", skill_registry_lifespan)
|
||||
@skill_registry.router.get("/stats")
|
||||
async def registry_stats():
|
||||
categories: dict[str, int] = {}
|
||||
for s in _cache.values():
|
||||
for s in P_CACHE.values():
|
||||
cat = s.get("category", "General")
|
||||
categories[cat] = categories.get(cat, 0) + 1
|
||||
return {
|
||||
"total": len(_cache),
|
||||
"total": len(P_CACHE),
|
||||
"categories": categories,
|
||||
"lastUpdated": _cache_updated_at,
|
||||
"lastUpdated": P_CACHE_UPDATED_AT,
|
||||
}
|
||||
|
||||
|
||||
@@ -159,7 +159,7 @@ async def registry_search(
|
||||
offset: int = Query(0, ge=0),
|
||||
category: str = Query("", description="Filter by category"),
|
||||
):
|
||||
pool = list(_cache.values())
|
||||
pool = list(P_CACHE.values())
|
||||
if category:
|
||||
cat_lower = category.lower()
|
||||
pool = [s for s in pool if s.get("category", "").lower() == cat_lower]
|
||||
@@ -192,7 +192,7 @@ async def registry_search(
|
||||
|
||||
@skill_registry.router.get("/detail/{skill_name:path}")
|
||||
async def registry_detail(skill_name: str):
|
||||
sk = _cache.get(skill_name)
|
||||
sk = P_CACHE.get(skill_name)
|
||||
if not sk:
|
||||
return {"error": "Skill not found"}, 404
|
||||
return {"skill": sk}
|
||||
|
||||
@@ -15,14 +15,14 @@ INDEX_PATH = os.path.join(SKILLS_DIR, ".skills_index.json")
|
||||
from backend.config.paths import SKILLS_WORKSPACE_DIR
|
||||
|
||||
|
||||
def _load_index() -> dict[str, dict]:
|
||||
def p_load_index() -> dict[str, dict]:
|
||||
if os.path.exists(INDEX_PATH):
|
||||
with open(INDEX_PATH) as f:
|
||||
return json.load(f)
|
||||
return {}
|
||||
|
||||
|
||||
def _save_index(index: dict[str, dict]):
|
||||
def p_save_index(index: dict[str, dict]):
|
||||
with open(INDEX_PATH, "w") as f:
|
||||
json.dump(index, f, indent=2)
|
||||
|
||||
@@ -32,7 +32,7 @@ def _save_index(index: dict[str, dict]):
|
||||
# `built_in: true` in the index. Users can edit the content (their
|
||||
# changes flow through to the matching agent's prompt on the next turn),
|
||||
# but they can't delete the file; the DELETE endpoint refuses with 409.
|
||||
def _built_in_skill_registry() -> list[dict]:
|
||||
def p_built_in_skill_registry() -> list[dict]:
|
||||
# Imported lazily so this module stays cheap to import from
|
||||
# everywhere (the skills outputs module pulls in pydantic+fastapi
|
||||
# transitively and we don't want a cycle).
|
||||
@@ -69,15 +69,15 @@ def _built_in_skill_registry() -> list[dict]:
|
||||
]
|
||||
|
||||
|
||||
def _seed_built_in_skills() -> None:
|
||||
def p_seed_built_in_skills() -> None:
|
||||
"""Copy each built-in skill into SKILLS_DIR if not already present, and
|
||||
ensure the index has the `built_in: true` flag so the UI and DELETE
|
||||
endpoint know to treat it specially. Idempotent; safe to call on
|
||||
every boot. Doesn't overwrite the file once it exists (so user edits
|
||||
are preserved across restarts and upgrades)."""
|
||||
index = _load_index()
|
||||
index = p_load_index()
|
||||
dirty = False
|
||||
for entry in _built_in_skill_registry():
|
||||
for entry in p_built_in_skill_registry():
|
||||
skill_id = entry["id"]
|
||||
fpath = os.path.join(SKILLS_DIR, f"{skill_id}.md")
|
||||
if not os.path.exists(fpath):
|
||||
@@ -103,7 +103,7 @@ def _seed_built_in_skills() -> None:
|
||||
index[skill_id] = meta
|
||||
dirty = True
|
||||
if dirty:
|
||||
_save_index(index)
|
||||
p_save_index(index)
|
||||
|
||||
|
||||
@asynccontextmanager
|
||||
@@ -111,7 +111,7 @@ async def skills_lifespan():
|
||||
os.makedirs(SKILLS_DIR, exist_ok=True)
|
||||
os.makedirs(SKILLS_WORKSPACE_DIR, exist_ok=True)
|
||||
try:
|
||||
_seed_built_in_skills()
|
||||
p_seed_built_in_skills()
|
||||
except Exception:
|
||||
# Don't block app startup on a skill-seed failure; the worst
|
||||
# case is the user has to manually paste the skill in once.
|
||||
@@ -122,9 +122,9 @@ async def skills_lifespan():
|
||||
skills = SubApp("skills", skills_lifespan)
|
||||
|
||||
|
||||
def _sync_skills() -> list[Skill]:
|
||||
def p_sync_skills() -> list[Skill]:
|
||||
"""Sync skills from the filesystem, updating the index."""
|
||||
index = _load_index()
|
||||
index = p_load_index()
|
||||
result = []
|
||||
|
||||
if os.path.exists(SKILLS_DIR):
|
||||
@@ -152,10 +152,10 @@ def _sync_skills() -> list[Skill]:
|
||||
|
||||
@skills.router.get("/list")
|
||||
async def list_skills():
|
||||
return {"skills": [s.model_dump() for s in _sync_skills()]}
|
||||
return {"skills": [s.model_dump() for s in p_sync_skills()]}
|
||||
|
||||
|
||||
def _parse_skill_frontmatter(raw: str) -> dict:
|
||||
def p_parse_skill_frontmatter(raw: str) -> dict:
|
||||
"""Extract YAML frontmatter fields from a SKILL.md file."""
|
||||
if not raw.startswith("---"):
|
||||
return {}
|
||||
@@ -207,7 +207,7 @@ async def read_skill_workspace(workspace_id: str):
|
||||
except json.JSONDecodeError:
|
||||
pass
|
||||
|
||||
frontmatter = _parse_skill_frontmatter(skill_content) if skill_content else {}
|
||||
frontmatter = p_parse_skill_frontmatter(skill_content) if skill_content else {}
|
||||
|
||||
return {
|
||||
"skill_content": skill_content,
|
||||
@@ -218,7 +218,7 @@ async def read_skill_workspace(workspace_id: str):
|
||||
|
||||
@skills.router.get("/{skill_id}")
|
||||
async def get_skill(skill_id: str):
|
||||
for s in _sync_skills():
|
||||
for s in p_sync_skills():
|
||||
if s.id == skill_id:
|
||||
return s.model_dump()
|
||||
raise HTTPException(status_code=404, detail="Skill not found")
|
||||
@@ -232,13 +232,13 @@ async def create_skill(body: SkillCreate):
|
||||
with open(fpath, "w") as f:
|
||||
f.write(body.content)
|
||||
|
||||
index = _load_index()
|
||||
index = p_load_index()
|
||||
index[slug] = {
|
||||
"name": body.name,
|
||||
"description": body.description,
|
||||
"command": body.command or slug,
|
||||
}
|
||||
_save_index(index)
|
||||
p_save_index(index)
|
||||
|
||||
skill = Skill(
|
||||
id=slug,
|
||||
@@ -262,7 +262,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
|
||||
with open(fpath, "w") as f:
|
||||
f.write(body.content)
|
||||
|
||||
index = _load_index()
|
||||
index = p_load_index()
|
||||
meta = index.get(skill_id, {})
|
||||
if body.name is not None:
|
||||
meta["name"] = body.name
|
||||
@@ -271,7 +271,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
|
||||
if body.command is not None:
|
||||
meta["command"] = body.command
|
||||
index[skill_id] = meta
|
||||
_save_index(index)
|
||||
p_save_index(index)
|
||||
|
||||
with open(fpath) as f:
|
||||
content = f.read()
|
||||
@@ -289,7 +289,7 @@ async def update_skill(skill_id: str, body: SkillUpdate):
|
||||
|
||||
@skills.router.delete("/{skill_id}")
|
||||
async def delete_skill(skill_id: str):
|
||||
index = _load_index()
|
||||
index = p_load_index()
|
||||
if index.get(skill_id, {}).get("built_in"):
|
||||
raise HTTPException(
|
||||
status_code=409,
|
||||
@@ -303,5 +303,5 @@ async def delete_skill(skill_id: str):
|
||||
if os.path.exists(fpath):
|
||||
os.remove(fpath)
|
||||
index.pop(skill_id, None)
|
||||
_save_index(index)
|
||||
p_save_index(index)
|
||||
return {"ok": True}
|
||||
|
||||
@@ -26,10 +26,10 @@ logger = logging.getLogger(__name__)
|
||||
|
||||
# Namespaces the hash so a raw hardware UUID never leaves the device. Public on
|
||||
# purpose (open-source): it only prevents transmitting the raw id, not a secret.
|
||||
_FP_SALT = "openswarm-free-trial-v1"
|
||||
P_FP_SALT = "openswarm-free-trial-v1"
|
||||
|
||||
|
||||
def _enabled() -> bool:
|
||||
def p_enabled() -> bool:
|
||||
# Default ON as of 1.2.80: the cloud free-trial proxy is live on prod
|
||||
# (api.openswarm.com) and arming + metered Haiku were verified end to end.
|
||||
# Set OPENSWARM_FREE_TRIAL_ENABLED=0 to force it off. The pool-shed gate +
|
||||
@@ -38,7 +38,7 @@ def _enabled() -> bool:
|
||||
return os.environ.get("OPENSWARM_FREE_TRIAL_ENABLED", "1") == "1"
|
||||
|
||||
|
||||
def _raw_hardware_id() -> str | None:
|
||||
def p_raw_hardware_id() -> str | None:
|
||||
"""A stable per-machine id that survives app reinstall / data wipe."""
|
||||
system = platform.system()
|
||||
try:
|
||||
@@ -67,17 +67,18 @@ def _raw_hardware_id() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _fingerprint(settings_obj) -> str | None:
|
||||
raw = _raw_hardware_id()
|
||||
def p_fingerprint(settings_obj) -> str | None:
|
||||
raw = p_raw_hardware_id()
|
||||
if not raw:
|
||||
# Fail-soft: installation_id is less durable (regenerates on wipe) but
|
||||
# better than nothing on a machine where the hardware id can't be read.
|
||||
raw = getattr(settings_obj, "installation_id", None)
|
||||
if not raw:
|
||||
return None
|
||||
return hashlib.sha256((_FP_SALT + raw).encode("utf-8")).hexdigest()
|
||||
return hashlib.sha256((P_FP_SALT + raw).encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
# Public - called by settings.py
|
||||
def has_own_model(s) -> bool:
|
||||
"""True if the user already has any real model path in settings; never shadow it."""
|
||||
if any(getattr(s, k, None) for k in (
|
||||
@@ -95,7 +96,7 @@ def has_own_model(s) -> bool:
|
||||
return False
|
||||
|
||||
|
||||
async def _has_connected_subscription() -> bool:
|
||||
async def p_has_connected_subscription() -> bool:
|
||||
"""True if 9Router holds a live Claude/ChatGPT/Gemini subscription. Those
|
||||
connections live in 9Router, not settings, so the sync check above misses
|
||||
them; this catches a sub connected while the trial was armed."""
|
||||
@@ -112,11 +113,11 @@ async def _has_connected_subscription() -> bool:
|
||||
return False
|
||||
|
||||
|
||||
def _proxy_base(settings_obj) -> str:
|
||||
def p_proxy_base(settings_obj) -> str:
|
||||
return (getattr(settings_obj, "openswarm_proxy_url", None) or OPENSWARM_DEFAULT_PROXY_URL).rstrip("/")
|
||||
|
||||
|
||||
async def _sync_routing(settings_obj) -> None:
|
||||
async def p_sync_routing(settings_obj) -> None:
|
||||
try:
|
||||
from backend.apps.nine_router.sync_custom import sync_pro_routing
|
||||
await sync_pro_routing(settings_obj)
|
||||
@@ -124,6 +125,7 @@ async def _sync_routing(settings_obj) -> None:
|
||||
logger.debug("free-trial routing sync skipped: %s", e)
|
||||
|
||||
|
||||
# Public - called by agent_manager.py
|
||||
async def clear_free_trial(settings_obj) -> None:
|
||||
"""Drop the trial token and revert to own_key. Keeps free_trial_remaining
|
||||
(so the UI knows it's spent) and never touches a real paid mode."""
|
||||
@@ -131,29 +133,30 @@ async def clear_free_trial(settings_obj) -> None:
|
||||
settings_obj.connection_mode = "own_key"
|
||||
settings_obj.free_trial_token = None
|
||||
await save_settings_async(settings_obj)
|
||||
await _sync_routing(settings_obj)
|
||||
await p_sync_routing(settings_obj)
|
||||
|
||||
|
||||
# Public - called by subscription.router.py
|
||||
async def arm_free_trial(settings_obj) -> dict:
|
||||
"""Mint (or re-fetch) the machine's grant and, if runs remain, flip into
|
||||
free-trial mode. Guarded: never arms over a real key/subscription."""
|
||||
if not _enabled():
|
||||
if not p_enabled():
|
||||
return {"armed": False, "reason": "disabled"}
|
||||
mode = getattr(settings_obj, "connection_mode", "own_key")
|
||||
if mode not in ("own_key", "free-trial"):
|
||||
return {"armed": False, "reason": "other_mode"}
|
||||
if has_own_model(settings_obj) or await _has_connected_subscription():
|
||||
if has_own_model(settings_obj) or await p_has_connected_subscription():
|
||||
# A real model exists now (key, custom provider, or a 9Router sub). If we
|
||||
# were on the free lane, hand the wheel back instead of re-arming.
|
||||
if mode == "free-trial":
|
||||
await clear_free_trial(settings_obj)
|
||||
return {"armed": False, "reason": "has_model"}
|
||||
|
||||
fp = _fingerprint(settings_obj)
|
||||
fp = p_fingerprint(settings_obj)
|
||||
if not fp:
|
||||
return {"armed": False, "reason": "no_fingerprint"}
|
||||
|
||||
base = _proxy_base(settings_obj)
|
||||
base = p_proxy_base(settings_obj)
|
||||
payload: dict = {"fingerprint_hash": fp}
|
||||
if getattr(settings_obj, "installation_id", None):
|
||||
payload["install_id"] = settings_obj.installation_id
|
||||
@@ -185,7 +188,7 @@ async def arm_free_trial(settings_obj) -> dict:
|
||||
if not cur.startswith(("sonnet", "haiku", "opus")):
|
||||
settings_obj.default_model = "haiku"
|
||||
await save_settings_async(settings_obj)
|
||||
await _sync_routing(settings_obj)
|
||||
await p_sync_routing(settings_obj)
|
||||
return {"armed": True, "runs_remaining": remaining, "runs_limit": settings_obj.free_trial_runs_limit}
|
||||
|
||||
# Already spent on this machine: record it but don't arm.
|
||||
@@ -193,6 +196,7 @@ async def arm_free_trial(settings_obj) -> dict:
|
||||
return {"armed": False, "reason": "exhausted", "runs_remaining": 0}
|
||||
|
||||
|
||||
# Public - called by subscription.router.py
|
||||
async def refresh_free_trial(settings_obj) -> dict:
|
||||
"""Re-read remaining runs from the cloud. Called after a session ends so the
|
||||
onboarding 'runs low' nudge stays honest. Clears the trial when spent."""
|
||||
@@ -200,7 +204,7 @@ async def refresh_free_trial(settings_obj) -> dict:
|
||||
if getattr(settings_obj, "connection_mode", "own_key") != "free-trial" or not token:
|
||||
return {"connected": False, "runs_remaining": getattr(settings_obj, "free_trial_remaining", None)}
|
||||
|
||||
base = _proxy_base(settings_obj)
|
||||
base = p_proxy_base(settings_obj)
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5.0) as client:
|
||||
r = await client.post(
|
||||
|
||||
@@ -26,7 +26,7 @@ async def subscription_lifespan():
|
||||
subscription = SubApp("subscription", subscription_lifespan)
|
||||
|
||||
|
||||
def _proxy_url() -> str:
|
||||
def p_proxy_url() -> str:
|
||||
"""Cloud router base URL. Overridable per-user via settings, falling back
|
||||
to the module-default. No trailing slash."""
|
||||
settings_obj = load_settings()
|
||||
@@ -35,7 +35,7 @@ def _proxy_url() -> str:
|
||||
return url.rstrip("/")
|
||||
|
||||
|
||||
async def _sync_pro_routing(settings_obj) -> None:
|
||||
async def p_sync_pro_routing(settings_obj) -> None:
|
||||
"""Mirror connection state into 9Router's Claude lane (WebSearch on
|
||||
non-Claude primaries). PUT /api/settings no longer carries these fields,
|
||||
so the state-change endpoints here are the only trigger left."""
|
||||
@@ -46,7 +46,7 @@ async def _sync_pro_routing(settings_obj) -> None:
|
||||
logger.debug("pro routing sync skipped: %s", e)
|
||||
|
||||
|
||||
async def _clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None:
|
||||
async def p_clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None:
|
||||
"""Revert to own_key mode and drop OpenSwarm Pro routing state.
|
||||
|
||||
`drop_bearer=True` (the default) is the original behavior, used when the
|
||||
@@ -66,17 +66,17 @@ async def _clear_subscription(settings_obj, *, drop_bearer: bool = True) -> None
|
||||
settings_obj.openswarm_subscription_expires = None
|
||||
settings_obj.openswarm_usage_cached = None
|
||||
await save_settings_async(settings_obj)
|
||||
_sync_subscription_identity(settings_obj)
|
||||
await _sync_pro_routing(settings_obj)
|
||||
p_sync_subscription_identity(settings_obj)
|
||||
await p_sync_pro_routing(settings_obj)
|
||||
|
||||
|
||||
def _sync_subscription_identity(settings_obj) -> None:
|
||||
def p_sync_subscription_identity(settings_obj) -> None:
|
||||
"""Push the installation's current subscription state into service-sync person
|
||||
properties so every event from this user is segmentable by plan /
|
||||
paying-vs-free. Safe to call from hot paths; service-sync is fire-and-forget
|
||||
and swallows errors internally."""
|
||||
try:
|
||||
from backend.apps.service.client import identify as _identify
|
||||
from backend.apps.service.client import identify
|
||||
except Exception:
|
||||
return
|
||||
mode = getattr(settings_obj, "connection_mode", "own_key")
|
||||
@@ -92,7 +92,7 @@ def _sync_subscription_identity(settings_obj) -> None:
|
||||
if is_paying and expires:
|
||||
props["subscription_expires"] = expires
|
||||
try:
|
||||
_identify(props)
|
||||
identify(props)
|
||||
except Exception as e:
|
||||
logger.debug("identify sync failed: %s", e)
|
||||
|
||||
@@ -118,7 +118,7 @@ async def activate(body: ActivateRequest):
|
||||
if not body.token or len(body.token) < 16:
|
||||
raise HTTPException(status_code=400, detail="Invalid token")
|
||||
|
||||
proxy = _proxy_url()
|
||||
proxy = p_proxy_url()
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
r = await client.get(
|
||||
@@ -165,8 +165,8 @@ async def activate(body: ActivateRequest):
|
||||
settings_obj.openswarm_usage_cached = usage
|
||||
|
||||
await save_settings_async(settings_obj)
|
||||
_sync_subscription_identity(settings_obj)
|
||||
await _sync_pro_routing(settings_obj)
|
||||
p_sync_subscription_identity(settings_obj)
|
||||
await p_sync_pro_routing(settings_obj)
|
||||
return {"ok": True, "plan": settings_obj.openswarm_subscription_plan}
|
||||
|
||||
|
||||
@@ -199,7 +199,7 @@ async def status():
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=5.0) as client:
|
||||
r = await client.get(
|
||||
f"{_proxy_url()}/api/me",
|
||||
f"{p_proxy_url()}/api/me",
|
||||
headers={"Authorization": f"Bearer {bearer}"},
|
||||
)
|
||||
upstream_code = r.status_code
|
||||
@@ -219,7 +219,7 @@ async def status():
|
||||
# through a dead subscription. Settings UI sees connected=False and
|
||||
# falls back to the Subscribe CTA; chat reverts to own_key routing.
|
||||
if upstream_code in (401, 402):
|
||||
await _clear_subscription(settings_obj)
|
||||
await p_clear_subscription(settings_obj)
|
||||
return {
|
||||
"connected": False,
|
||||
"connection_mode": "own_key",
|
||||
@@ -253,34 +253,34 @@ async def sync():
|
||||
already had."""
|
||||
# Lazy-import the service-sync helper so subscription/router doesn't pay the
|
||||
# cost when analytics are disabled.
|
||||
from backend.apps.service.client import sync as _sync
|
||||
from backend.apps.service.client import sync as client_sync
|
||||
|
||||
settings_obj = load_settings()
|
||||
bearer = getattr(settings_obj, "openswarm_bearer_token", None)
|
||||
mode = getattr(settings_obj, "connection_mode", "own_key")
|
||||
|
||||
if mode != "openswarm-pro" or not bearer:
|
||||
_sync(settings_obj.model_dump())
|
||||
client_sync(settings_obj.model_dump())
|
||||
return {"ok": True, "synced": False, "connection_mode": mode}
|
||||
|
||||
try:
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
r = await client.post(
|
||||
f"{_proxy_url()}/api/subscription/sync",
|
||||
f"{p_proxy_url()}/api/subscription/sync",
|
||||
headers={"Authorization": f"Bearer {bearer}"},
|
||||
)
|
||||
except httpx.HTTPError as e:
|
||||
logger.debug("subscription/sync live fetch failed: %s", e)
|
||||
_sync(settings_obj.model_dump())
|
||||
client_sync(settings_obj.model_dump())
|
||||
return {"ok": True, "synced": False, "reason": "network"}
|
||||
|
||||
# Same 401/402 handling as /status: if Stripe-side reconciliation proves
|
||||
# the bearer is dead or the sub expired, clear local state so the app
|
||||
# reverts to own_key instead of hammering a useless token.
|
||||
if r.status_code in (401, 402):
|
||||
await _clear_subscription(settings_obj)
|
||||
await p_clear_subscription(settings_obj)
|
||||
reason = "revoked" if r.status_code == 401 else "expired"
|
||||
_sync(settings_obj.model_dump())
|
||||
client_sync(settings_obj.model_dump())
|
||||
return {
|
||||
"ok": True,
|
||||
"synced": False,
|
||||
@@ -290,7 +290,7 @@ async def sync():
|
||||
|
||||
if r.status_code != 200:
|
||||
logger.debug("subscription/sync got %s from cloud: %s", r.status_code, r.text[:200])
|
||||
_sync(settings_obj.model_dump())
|
||||
client_sync(settings_obj.model_dump())
|
||||
return {"ok": True, "synced": False, "reason": "upstream"}
|
||||
|
||||
data = r.json()
|
||||
@@ -307,8 +307,8 @@ async def sync():
|
||||
datetime.fromtimestamp(period_end_ms / 1000, tz=timezone.utc).isoformat()
|
||||
)
|
||||
await save_settings_async(settings_obj)
|
||||
_sync_subscription_identity(settings_obj)
|
||||
_sync(settings_obj.model_dump())
|
||||
p_sync_subscription_identity(settings_obj)
|
||||
client_sync(settings_obj.model_dump())
|
||||
return {
|
||||
"ok": True,
|
||||
"synced": bool(data.get("synced")),
|
||||
@@ -333,7 +333,7 @@ async def portal():
|
||||
|
||||
async with httpx.AsyncClient(timeout=10.0) as client:
|
||||
r = await client.post(
|
||||
f"{_proxy_url()}/api/billing/portal",
|
||||
f"{p_proxy_url()}/api/billing/portal",
|
||||
headers={"Authorization": f"Bearer {bearer}"},
|
||||
)
|
||||
if r.status_code >= 400:
|
||||
@@ -373,5 +373,5 @@ async def disconnect():
|
||||
does NOT sign the user out of OpenSwarm (use /api/auth/signout for that).
|
||||
Useful when a user wants to temporarily route through their own API key
|
||||
without losing their account state."""
|
||||
await _clear_subscription(load_settings(), drop_bearer=False)
|
||||
await p_clear_subscription(load_settings(), drop_bearer=False)
|
||||
return {"ok": True}
|
||||
|
||||
@@ -11,18 +11,19 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
# Public - called by main.py, agent_manager.py, prompt_context.py
|
||||
def sanitize_mcp_server_name(name: str) -> str:
|
||||
"""Convert a tool name into a valid MCP server identifier (alphanumeric + hyphens)."""
|
||||
return re.sub(r"[^a-z0-9]+", "-", name.lower()).strip("-")
|
||||
|
||||
|
||||
def _extra_bin_dirs() -> list[str]:
|
||||
def p_extra_bin_dirs() -> list[str]:
|
||||
"""Well-known user-local bin directories that may not be on PATH in packaged apps."""
|
||||
home = os.path.expanduser("~")
|
||||
# Bundled uv-bin (ships uvx for non-dev users)
|
||||
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
dirs = [
|
||||
os.path.join(_backend, "uv-bin"),
|
||||
os.path.join(backend_dir, "uv-bin"),
|
||||
os.path.join(home, ".bun", "bin"),
|
||||
os.path.join(home, ".cargo", "bin"),
|
||||
os.path.join(home, ".local", "bin"),
|
||||
@@ -46,7 +47,8 @@ def _extra_bin_dirs() -> list[str]:
|
||||
return dirs
|
||||
|
||||
|
||||
def _resolve_command(command: str) -> str | None:
|
||||
# Public - called by mcp_discovery.py
|
||||
def resolve_command(command: str) -> str | None:
|
||||
"""Find a command on PATH, falling back to common user-local bin directories
|
||||
and bundled binaries (uv-bin for uvx/uv)."""
|
||||
found = shutil.which(command)
|
||||
@@ -59,24 +61,25 @@ def _resolve_command(command: str) -> str | None:
|
||||
suffixes = [""] + os.environ.get("PATHEXT", ".COM;.EXE;.BAT;.CMD").lower().split(os.pathsep)
|
||||
else:
|
||||
suffixes = [""]
|
||||
def _probe(directory: str) -> str | None:
|
||||
def probe(directory: str) -> str | None:
|
||||
for suffix in suffixes:
|
||||
candidate = os.path.join(directory, command + suffix)
|
||||
if os.path.isfile(candidate) and os.access(candidate, os.X_OK):
|
||||
return candidate
|
||||
return None
|
||||
for d in _extra_bin_dirs():
|
||||
hit = _probe(d)
|
||||
for d in p_extra_bin_dirs():
|
||||
hit = probe(d)
|
||||
if hit:
|
||||
return hit
|
||||
# Check bundled uv-bin directory (ships uv/uvx for non-dev users)
|
||||
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
return _probe(os.path.join(_backend, "uv-bin"))
|
||||
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
return probe(os.path.join(backend_dir, "uv-bin"))
|
||||
|
||||
|
||||
def _augmented_path() -> str:
|
||||
# Public - called by mcp_discovery.py
|
||||
def augmented_path() -> str:
|
||||
"""Return PATH with extra bin dirs prepended (for child process environments)."""
|
||||
extra = [d for d in _extra_bin_dirs() if os.path.isdir(d)]
|
||||
extra = [d for d in p_extra_bin_dirs() if os.path.isdir(d)]
|
||||
current = os.environ.get("PATH", "")
|
||||
seen: set[str] = set()
|
||||
parts: list[str] = []
|
||||
@@ -87,6 +90,7 @@ def _augmented_path() -> str:
|
||||
return os.pathsep.join(parts)
|
||||
|
||||
|
||||
# Public - called by tools_lib.py
|
||||
def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
"""Build the claude_agent_sdk mcp_servers config entry for a tool.
|
||||
|
||||
@@ -130,9 +134,9 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
# proxy that forwards the refresh to our cloud's pool-aware
|
||||
# /api/oauth/google/refresh endpoint; CLIENT_ID/SECRET become
|
||||
# unused placeholders (gauth.py only validates non-empty).
|
||||
_port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
env["GOOGLE_WORKSPACE_TOKEN_URI"] = (
|
||||
f"http://127.0.0.1:{_port}/api/tools/google-oauth-token"
|
||||
f"http://127.0.0.1:{port}/api/tools/google-oauth-token"
|
||||
)
|
||||
env.setdefault("GOOGLE_WORKSPACE_CLIENT_ID", "openswarm-proxy")
|
||||
env.setdefault("GOOGLE_WORKSPACE_CLIENT_SECRET", "openswarm-proxy")
|
||||
@@ -167,9 +171,9 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
# The shim runs as a subprocess and needs to import
|
||||
# `backend.apps.discord_mcp_shim`; set PYTHONPATH to the project
|
||||
# root (parent of the backend/ dir) so that import resolves.
|
||||
_project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
project_root = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
existing_pp = env.get("PYTHONPATH") or os.environ.get("PYTHONPATH", "")
|
||||
env["PYTHONPATH"] = (_project_root + os.pathsep + existing_pp) if existing_pp else _project_root
|
||||
env["PYTHONPATH"] = (project_root + os.pathsep + existing_pp) if existing_pp else project_root
|
||||
|
||||
# Microsoft 365 MCP: use a stable token cache path shared across process spawns
|
||||
if tool.name.lower() == "microsoft 365" and config.get("type") == "stdio":
|
||||
@@ -196,7 +200,7 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
if config["command"] in ("npx", "bunx"):
|
||||
pkg_name = next((a for a in (config.get("args") or []) if not a.startswith("-")), None)
|
||||
if pkg_name:
|
||||
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
electron_path = os.environ.get("OPENSWARM_ELECTRON_PATH")
|
||||
# Two bundle layouts in mcp-bundles/, checked in priority order:
|
||||
#
|
||||
@@ -217,8 +221,8 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
# Scoped names get flattened ("@softeria/ms-365-mcp-server"
|
||||
# -> "softeria-ms-365-mcp-server") for filesystem safety.
|
||||
safe_bundle = pkg_name.replace("/", "-").replace("@", "")
|
||||
bundle_dir_path = os.path.join(_backend, "mcp-bundles", safe_bundle, "dist", "index.js")
|
||||
bundle_file_path = os.path.join(_backend, "mcp-bundles", f"{safe_bundle}.js")
|
||||
bundle_dir_path = os.path.join(backend_dir, "mcp-bundles", safe_bundle, "dist", "index.js")
|
||||
bundle_file_path = os.path.join(backend_dir, "mcp-bundles", f"{safe_bundle}.js")
|
||||
bundle_path = None
|
||||
if os.path.isfile(bundle_dir_path):
|
||||
bundle_path = bundle_dir_path
|
||||
@@ -242,12 +246,12 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
else:
|
||||
# Check for pre-installed npm package (works in both dev and packaged modes)
|
||||
safe_dir = pkg_name.replace("/", "-").replace("@", "")
|
||||
npm_dir = os.path.join(_backend, "npm-servers", safe_dir)
|
||||
npm_dir = os.path.join(backend_dir, "npm-servers", safe_dir)
|
||||
pkg_json_path = os.path.join(npm_dir, "node_modules", pkg_name, "package.json")
|
||||
if os.path.isfile(pkg_json_path):
|
||||
import json as _json
|
||||
import json
|
||||
with open(pkg_json_path) as f:
|
||||
pkg_meta = _json.load(f)
|
||||
pkg_meta = json.load(f)
|
||||
bin_field = pkg_meta.get("bin", {})
|
||||
entry = list(bin_field.values())[0] if isinstance(bin_field, dict) else bin_field
|
||||
# Same priority as 9Router / MCP-bundle paths: bundled node > system node > Electron-as-Node.
|
||||
@@ -262,33 +266,33 @@ def derive_mcp_config(tool: ToolDefinition) -> Optional[dict]:
|
||||
logger.info(f"Using pre-installed npm MCP server for {pkg_name}")
|
||||
|
||||
if not os.path.isabs(config.get("command", "")):
|
||||
resolved = _resolve_command(config["command"])
|
||||
resolved = resolve_command(config["command"])
|
||||
if resolved:
|
||||
config["command"] = resolved
|
||||
else:
|
||||
logger.warning(f"Command '{config['command']}' not found on PATH or bundled directories")
|
||||
env = config.setdefault("env", {})
|
||||
env.setdefault("PATH", _augmented_path())
|
||||
env.setdefault("PATH", augmented_path())
|
||||
env.setdefault("PYTHONPATH", "")
|
||||
# Point uv/uvx at our bundled Python; avoids macOS CLT popup on fresh Macs
|
||||
# and avoids downloading Python at runtime
|
||||
_is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
|
||||
_is_windows = sys.platform == "win32"
|
||||
if _is_packaged:
|
||||
_resources = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
if _is_windows:
|
||||
_bundled_python = os.path.join(_resources, "python-env", "python.exe")
|
||||
is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
|
||||
is_windows = sys.platform == "win32"
|
||||
if is_packaged:
|
||||
resources_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__)))))
|
||||
if is_windows:
|
||||
bundled_python = os.path.join(resources_dir, "python-env", "python.exe")
|
||||
else:
|
||||
_bundled_python = os.path.join(_resources, "python-env", "bin", "python3")
|
||||
if os.path.exists(_bundled_python):
|
||||
env.setdefault("UV_PYTHON", _bundled_python)
|
||||
bundled_python = os.path.join(resources_dir, "python-env", "bin", "python3")
|
||||
if os.path.exists(bundled_python):
|
||||
env.setdefault("UV_PYTHON", bundled_python)
|
||||
else:
|
||||
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
if _is_windows:
|
||||
_venv_python = os.path.join(_backend, ".venv", "Scripts", "python.exe")
|
||||
backend_dir = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
if is_windows:
|
||||
venv_python = os.path.join(backend_dir, ".venv", "Scripts", "python.exe")
|
||||
else:
|
||||
_venv_python = os.path.join(_backend, ".venv", "bin", "python3")
|
||||
if os.path.exists(_venv_python):
|
||||
env.setdefault("UV_PYTHON", _venv_python)
|
||||
venv_python = os.path.join(backend_dir, ".venv", "bin", "python3")
|
||||
if os.path.exists(venv_python):
|
||||
env.setdefault("UV_PYTHON", venv_python)
|
||||
|
||||
return config
|
||||
|
||||
@@ -8,12 +8,12 @@ import shutil
|
||||
import httpx
|
||||
from fastapi import HTTPException
|
||||
|
||||
from backend.apps.tools_lib.mcp_config import _augmented_path, _resolve_command
|
||||
from backend.apps.tools_lib.mcp_config import augmented_path, resolve_command
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _parse_sse_json(text: str) -> dict | None:
|
||||
def p_parse_sse_json(text: str) -> dict | None:
|
||||
"""Extract JSON from an SSE response body (handles `data: {...}` lines)."""
|
||||
for line in text.splitlines():
|
||||
stripped = line.strip()
|
||||
@@ -30,7 +30,8 @@ def _parse_sse_json(text: str) -> dict | None:
|
||||
return None
|
||||
|
||||
|
||||
async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> list[dict]:
|
||||
# Public - called by tools_lib.py
|
||||
async def discover_mcp_tools_http(url: str, headers: dict | None = None) -> list[dict]:
|
||||
"""Connect to a Streamable HTTP MCP server and call tools/list via JSON-RPC POST."""
|
||||
h = {
|
||||
"Content-Type": "application/json",
|
||||
@@ -62,7 +63,7 @@ async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> lis
|
||||
|
||||
ct = list_resp.headers.get("content-type", "")
|
||||
if "text/event-stream" in ct:
|
||||
data = _parse_sse_json(list_resp.text)
|
||||
data = p_parse_sse_json(list_resp.text)
|
||||
else:
|
||||
data = list_resp.json()
|
||||
|
||||
@@ -73,7 +74,8 @@ async def _discover_mcp_tools_http(url: str, headers: dict | None = None) -> lis
|
||||
return [{"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} for t in tools_list]
|
||||
|
||||
|
||||
async def _discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list[dict]:
|
||||
# Public - called by tools_lib.py
|
||||
async def discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list[dict]:
|
||||
"""Connect to a legacy SSE MCP server (GET event-stream + POST messages) and call tools/list."""
|
||||
from mcp.client.sse import sse_client
|
||||
from mcp import ClientSession
|
||||
@@ -99,10 +101,10 @@ async def _discover_mcp_tools_sse(url: str, headers: dict | None = None) -> list
|
||||
raise HTTPException(status_code=502, detail=f"SSE discovery failed: {first}") from first
|
||||
|
||||
|
||||
_NPX_CACHE_RE = re.compile(r"_npx[/\\]([0-9a-f]{8,})[/\\]")
|
||||
P_NPX_CACHE_RE = re.compile(r"_npx[/\\]([0-9a-f]{8,})[/\\]")
|
||||
|
||||
|
||||
def _try_heal_npx_cache(stderr: str) -> str | None:
|
||||
def p_try_heal_npx_cache(stderr: str) -> str | None:
|
||||
"""On `ERR_MODULE_NOT_FOUND` pointing into `~/.npm/_npx/<hash>/`, wipe that one dir.
|
||||
|
||||
Why: interrupted npx installs leave a `package-lock.json` in the cache dir so
|
||||
@@ -112,7 +114,7 @@ def _try_heal_npx_cache(stderr: str) -> str | None:
|
||||
"""
|
||||
if "ERR_MODULE_NOT_FOUND" not in stderr:
|
||||
return None
|
||||
m = _NPX_CACHE_RE.search(stderr)
|
||||
m = P_NPX_CACHE_RE.search(stderr)
|
||||
if not m:
|
||||
return None
|
||||
hash_ = m.group(1)
|
||||
@@ -124,18 +126,19 @@ def _try_heal_npx_cache(stderr: str) -> str | None:
|
||||
return hash_
|
||||
|
||||
|
||||
async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None, env: dict | None = None, _attempt: int = 0) -> list[dict]:
|
||||
# Public - called by tools_lib.py
|
||||
async def discover_mcp_tools_stdio(command: str, args: list[str] | None = None, env: dict | None = None, attempt: int = 0) -> list[dict]:
|
||||
"""Spawn a stdio MCP server process and call tools/list via JSON-RPC over stdin/stdout.
|
||||
|
||||
On the first attempt, a failure that looks like corrupted npx cache
|
||||
(`ERR_MODULE_NOT_FOUND` pointing into `~/.npm/_npx/<hash>/`) triggers one
|
||||
auto-heal + retry. No heal on `_attempt >= 1`.
|
||||
auto-heal + retry. No heal on `attempt >= 1`.
|
||||
"""
|
||||
cmd_path = _resolve_command(command)
|
||||
cmd_path = resolve_command(command)
|
||||
if not cmd_path:
|
||||
raise HTTPException(status_code=400, detail=f"Command '{command}' not found on PATH or common install locations")
|
||||
|
||||
proc_env = {**os.environ, **(env or {}), "PATH": _augmented_path()}
|
||||
proc_env = {**os.environ, **(env or {}), "PATH": augmented_path()}
|
||||
proc_env.pop("PYTHONPATH", None)
|
||||
|
||||
proc = await asyncio.create_subprocess_exec(
|
||||
@@ -155,7 +158,7 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
|
||||
# the opaque "discovery failed" we used to show.
|
||||
stderr_tail: list[str] = []
|
||||
|
||||
async def _drain_stderr() -> None:
|
||||
async def drain_stderr() -> None:
|
||||
try:
|
||||
while True:
|
||||
chunk = await proc.stderr.readline()
|
||||
@@ -169,14 +172,14 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
|
||||
except Exception:
|
||||
return
|
||||
|
||||
stderr_task = asyncio.create_task(_drain_stderr())
|
||||
stderr_task = asyncio.create_task(drain_stderr())
|
||||
|
||||
async def _send(msg: dict) -> None:
|
||||
async def send(msg: dict) -> None:
|
||||
line = json.dumps(msg) + "\n"
|
||||
proc.stdin.write(line.encode())
|
||||
await proc.stdin.drain()
|
||||
|
||||
async def _recv(timeout_s: float = 30.0) -> dict:
|
||||
async def recv(timeout_s: float = 30.0) -> dict:
|
||||
"""Read JSON-RPC responses, skipping notification lines (no 'id' field)."""
|
||||
while True:
|
||||
line = await asyncio.wait_for(proc.stdout.readline(), timeout=timeout_s)
|
||||
@@ -204,7 +207,7 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
|
||||
return data
|
||||
|
||||
try:
|
||||
await _send({
|
||||
await send({
|
||||
"jsonrpc": "2.0", "id": 1, "method": "initialize",
|
||||
"params": {
|
||||
"protocolVersion": "2025-03-26",
|
||||
@@ -217,12 +220,12 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
|
||||
# AV-scan every file npm writes; total install time often exceeds
|
||||
# 60 s and occasionally pushes past 90 s. Subsequent reads run
|
||||
# against an already-running server and stay at the default 30 s.
|
||||
await _recv(timeout_s=120.0)
|
||||
await recv(timeout_s=120.0)
|
||||
|
||||
await _send({"jsonrpc": "2.0", "method": "notifications/initialized"})
|
||||
await send({"jsonrpc": "2.0", "method": "notifications/initialized"})
|
||||
|
||||
await _send({"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}})
|
||||
data = await _recv()
|
||||
await send({"jsonrpc": "2.0", "id": 2, "method": "tools/list", "params": {}})
|
||||
data = await recv()
|
||||
|
||||
tools_list = data.get("result", {}).get("tools", [])
|
||||
return [{"name": t.get("name", ""), "description": t.get("description", ""), "inputSchema": t.get("inputSchema")} for t in tools_list]
|
||||
@@ -231,8 +234,8 @@ async def _discover_mcp_tools_stdio(command: str, args: list[str] | None = None,
|
||||
# Heal-on-corrupt-npx-cache still triggers from the EOF branch,
|
||||
# which now includes the full stderr tail in `e.detail`; so the
|
||||
# ERR_MODULE_NOT_FOUND signature is still discoverable here.
|
||||
if _attempt == 0 and _try_heal_npx_cache(str(e.detail) if e.detail is not None else ""):
|
||||
return await _discover_mcp_tools_stdio(command, args, env, _attempt=1)
|
||||
if attempt == 0 and p_try_heal_npx_cache(str(e.detail) if e.detail is not None else ""):
|
||||
return await discover_mcp_tools_stdio(command, args, env, attempt=1)
|
||||
raise
|
||||
except asyncio.TimeoutError:
|
||||
# Most common cause: cold npx cache on Windows. The npm install
|
||||
|
||||
@@ -13,7 +13,7 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
|
||||
def _save(tool: ToolDefinition) -> None:
|
||||
def p_save(tool: ToolDefinition) -> None:
|
||||
with open(os.path.join(DATA_DIR, f"{tool.id}.json"), "w") as f:
|
||||
json.dump(tool.model_dump(), f, indent=2)
|
||||
|
||||
@@ -22,7 +22,7 @@ def _save(tool: ToolDefinition) -> None:
|
||||
# through the Fly cloud-proxy so client_secret values never ship inside the
|
||||
# desktop binary. v1.0.28 was the last release that used a local Google
|
||||
# callback with the client_secret in backend/.env.
|
||||
_TOOL_NAME_TO_PROVIDER = {
|
||||
P_TOOL_NAME_TO_PROVIDER = {
|
||||
"airtable": "airtable",
|
||||
"hubspot": "hubspot",
|
||||
"discord": "discord",
|
||||
@@ -34,11 +34,13 @@ _TOOL_NAME_TO_PROVIDER = {
|
||||
}
|
||||
|
||||
|
||||
def _proxied_provider_for(tool: ToolDefinition) -> Optional[str]:
|
||||
return _TOOL_NAME_TO_PROVIDER.get(tool.name.lower())
|
||||
# Public - called by tools_lib.py
|
||||
def proxied_provider_for(tool: ToolDefinition) -> Optional[str]:
|
||||
return P_TOOL_NAME_TO_PROVIDER.get(tool.name.lower())
|
||||
|
||||
|
||||
def _persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
|
||||
# Public - called by tools_lib.py
|
||||
def persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
|
||||
"""Normalise the cloud's claim response into tool.oauth_tokens.
|
||||
|
||||
Per-provider shaping mirrors what the v1.0.25 local-callback flow used
|
||||
@@ -77,7 +79,7 @@ def _persist_cloud_tokens(tool: ToolDefinition, tokens: dict) -> None:
|
||||
tool.auth_status = "connected"
|
||||
|
||||
|
||||
async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry: int) -> Optional[str]:
|
||||
async def p_refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry: int) -> Optional[str]:
|
||||
"""Refresh an OAuth access_token by POSTing the refresh_token to the
|
||||
helper service. Per-provider wrappers below pass a default expires_in
|
||||
fallback for providers that don't return one.
|
||||
@@ -101,7 +103,7 @@ async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry
|
||||
# Provider rejected; user revoked at the provider's side. Mark
|
||||
# as needing re-auth so the UI prompts a Reconnect.
|
||||
tool.auth_status = "expired"
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
logger.warning(f"{provider} refresh rejected (user revoked); marking tool as expired")
|
||||
return None
|
||||
if resp.status_code != 200:
|
||||
@@ -121,13 +123,14 @@ async def _refresh_via_proxy(provider: str, tool: ToolDefinition, default_expiry
|
||||
# Backfill identity label on first successful refresh after upgrade.
|
||||
if not tool.connected_account_email and data.get("email"):
|
||||
tool.connected_account_email = data["email"]
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return new_token
|
||||
except Exception as e:
|
||||
logger.warning(f"{provider} cloud refresh exception for tool {tool.id}: {e}")
|
||||
return None
|
||||
|
||||
|
||||
# Public - called by tools_lib.py
|
||||
async def refresh_google_token(tool: ToolDefinition) -> Optional[str]:
|
||||
"""Refresh an expired Google access_token via the Fly cloud-proxy.
|
||||
|
||||
@@ -135,20 +138,22 @@ async def refresh_google_token(tool: ToolDefinition) -> Optional[str]:
|
||||
refresh_token. Same pattern as Airtable/HubSpot. Pre-v1.0.29 builds
|
||||
held the secret in their bundled .env; v1.0.29 removed it.
|
||||
"""
|
||||
return await _refresh_via_proxy("google", tool, default_expiry=3600)
|
||||
return await p_refresh_via_proxy("google", tool, default_expiry=3600)
|
||||
|
||||
|
||||
# Public - called by tools_lib.py
|
||||
async def refresh_airtable_token(tool: ToolDefinition) -> Optional[str]:
|
||||
"""Refresh an expired Airtable OAuth access_token."""
|
||||
return await _refresh_via_proxy("airtable", tool, default_expiry=7200)
|
||||
return await p_refresh_via_proxy("airtable", tool, default_expiry=7200)
|
||||
|
||||
|
||||
# Public - called by tools_lib.py
|
||||
async def refresh_hubspot_token(tool: ToolDefinition) -> Optional[str]:
|
||||
"""Refresh an expired HubSpot OAuth access_token."""
|
||||
return await _refresh_via_proxy("hubspot", tool, default_expiry=1800)
|
||||
return await p_refresh_via_proxy("hubspot", tool, default_expiry=1800)
|
||||
|
||||
|
||||
def _m365_server_script() -> str:
|
||||
# Public - called by tools_lib.py
|
||||
def m365_server_script() -> str:
|
||||
"""Return the on-disk path to the bundled MS365 MCP server entry.
|
||||
|
||||
v1.0.26 replaced the heavy backend/npm-servers/softeria-ms-365-mcp-server/
|
||||
@@ -158,9 +163,9 @@ def _m365_server_script() -> str:
|
||||
package.json) because cli.js reads __dirname/../package.json for the
|
||||
--version flag; see scripts/build-app.sh `build_mcp_bundle_dir`.
|
||||
"""
|
||||
_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
p_backend = os.path.dirname(os.path.dirname(os.path.dirname(os.path.abspath(__file__))))
|
||||
bundle = os.path.join(
|
||||
_backend, "mcp-bundles", "softeria-ms-365-mcp-server", "dist", "index.js",
|
||||
p_backend, "mcp-bundles", "softeria-ms-365-mcp-server", "dist", "index.js",
|
||||
)
|
||||
if os.path.isfile(bundle):
|
||||
return bundle
|
||||
@@ -168,12 +173,13 @@ def _m365_server_script() -> str:
|
||||
# was left over from before the bundle migration. Will return the legacy
|
||||
# path; if that doesn't exist either, the caller raises a clear error.
|
||||
return os.path.join(
|
||||
_backend, "npm-servers", "softeria-ms-365-mcp-server",
|
||||
p_backend, "npm-servers", "softeria-ms-365-mcp-server",
|
||||
"node_modules", "@softeria", "ms-365-mcp-server", "dist", "index.js",
|
||||
)
|
||||
|
||||
|
||||
def _m365_cache_env() -> dict[str, str]:
|
||||
# Public - called by tools_lib.py
|
||||
def m365_cache_env() -> dict[str, str]:
|
||||
cache_dir = os.path.join(os.path.expanduser("~"), ".openswarm")
|
||||
os.makedirs(cache_dir, exist_ok=True)
|
||||
return {
|
||||
|
||||
@@ -1,8 +1,8 @@
|
||||
_READ_PREFIXES = ("get", "list", "read", "search", "fetch", "find", "query", "count", "check", "describe", "show", "download", "browse", "analy", "explain")
|
||||
_WRITE_PREFIXES = ("create", "write", "delete", "update", "send", "remove", "modify", "add", "set", "put", "post", "patch", "insert", "move", "copy", "rename", "archive", "trash", "publish", "approve", "reject")
|
||||
P_READ_PREFIXES = ("get", "list", "read", "search", "fetch", "find", "query", "count", "check", "describe", "show", "download", "browse", "analy", "explain")
|
||||
P_WRITE_PREFIXES = ("create", "write", "delete", "update", "send", "remove", "modify", "add", "set", "put", "post", "patch", "insert", "move", "copy", "rename", "archive", "trash", "publish", "approve", "reject")
|
||||
|
||||
|
||||
_SERVICE_RULES: list[tuple[list[str], str, str]] = [
|
||||
P_SERVICE_RULES: list[tuple[list[str], str, str]] = [
|
||||
# (keywords, service_name, group)
|
||||
# Google Workspace
|
||||
(["gmail"], "Gmail", "Google"),
|
||||
@@ -31,20 +31,20 @@ _SERVICE_RULES: list[tuple[list[str], str, str]] = [
|
||||
]
|
||||
|
||||
|
||||
def _categorize_tool(name: str) -> str:
|
||||
def p_categorize_tool(name: str) -> str:
|
||||
lower = name.lower().replace("_", " ").replace("-", " ").strip()
|
||||
for word in lower.split():
|
||||
for prefix in _READ_PREFIXES:
|
||||
for prefix in P_READ_PREFIXES:
|
||||
if word.startswith(prefix):
|
||||
return "read"
|
||||
for prefix in _WRITE_PREFIXES:
|
||||
for prefix in P_WRITE_PREFIXES:
|
||||
if word.startswith(prefix):
|
||||
return "write"
|
||||
return "write"
|
||||
|
||||
|
||||
def _integration_domain(integration: str) -> str:
|
||||
"""Which curated _SERVICE_RULES set applies to this integration, if any. The Google rules use
|
||||
def p_integration_domain(integration: str) -> str:
|
||||
"""Which curated P_SERVICE_RULES set applies to this integration, if any. The Google rules use
|
||||
generic words (message/table/page/doc/script) that otherwise mis-tag Slack/Notion/Airtable/M365."""
|
||||
n = (integration or "").lower()
|
||||
if "google" in n:
|
||||
@@ -56,13 +56,13 @@ def _integration_domain(integration: str) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
def _extract_service(name: str, integration: str) -> tuple[str, str]:
|
||||
def p_extract_service(name: str, integration: str) -> tuple[str, str]:
|
||||
"""Map a tool name to (service, group). Curated rulesets apply only to the integration they were
|
||||
written for; every other integration groups under its own name so it isn't mislabeled as Google."""
|
||||
domain = _integration_domain(integration)
|
||||
domain = p_integration_domain(integration)
|
||||
if domain:
|
||||
lower = name.lower()
|
||||
for keywords, display, group in _SERVICE_RULES:
|
||||
for keywords, display, group in P_SERVICE_RULES:
|
||||
if group != domain:
|
||||
continue
|
||||
for kw in keywords:
|
||||
@@ -73,15 +73,16 @@ def _extract_service(name: str, integration: str) -> tuple[str, str]:
|
||||
return (integration or "Other"), ""
|
||||
|
||||
|
||||
def _classify_services(
|
||||
# Public - called by tools_lib.py
|
||||
def classify_services(
|
||||
tool_names: list[str], integration: str
|
||||
) -> tuple[dict[str, dict[str, list[str]]], dict[str, list[str]], list[str], list[str]]:
|
||||
"""Bucket tool names into services + service groups + read/write categories for one integration."""
|
||||
services: dict[str, dict[str, list[str]]] = {}
|
||||
service_groups: dict[str, list[str]] = {}
|
||||
for name in tool_names:
|
||||
cat = _categorize_tool(name)
|
||||
svc, group = _extract_service(name, integration)
|
||||
cat = p_categorize_tool(name)
|
||||
svc, group = p_extract_service(name, integration)
|
||||
services.setdefault(svc, {"read": [], "write": []})
|
||||
services[svc][cat].append(name)
|
||||
if group:
|
||||
|
||||
@@ -21,20 +21,20 @@ from backend.apps.tools_lib.oauth_config import OPENSWARM_OAUTH_BASE_URL
|
||||
# derive_mcp_config re-exported for agent_manager/main.
|
||||
from backend.apps.tools_lib.mcp_config import derive_mcp_config
|
||||
from backend.apps.tools_lib.mcp_discovery import (
|
||||
_discover_mcp_tools_http,
|
||||
_discover_mcp_tools_sse,
|
||||
_discover_mcp_tools_stdio,
|
||||
discover_mcp_tools_http,
|
||||
discover_mcp_tools_sse,
|
||||
discover_mcp_tools_stdio,
|
||||
)
|
||||
from backend.apps.tools_lib.tool_taxonomy import _classify_services
|
||||
from backend.apps.tools_lib.tool_taxonomy import classify_services
|
||||
# refresh_* re-exported for agent_manager.
|
||||
from backend.apps.tools_lib.oauth_tokens import (
|
||||
_proxied_provider_for,
|
||||
_persist_cloud_tokens,
|
||||
proxied_provider_for,
|
||||
persist_cloud_tokens,
|
||||
refresh_google_token,
|
||||
refresh_airtable_token,
|
||||
refresh_hubspot_token,
|
||||
_m365_server_script,
|
||||
_m365_cache_env,
|
||||
m365_server_script,
|
||||
m365_cache_env,
|
||||
)
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
@@ -43,8 +43,8 @@ logger = logging.getLogger(__name__)
|
||||
@asynccontextmanager
|
||||
async def tools_lib_lifespan():
|
||||
os.makedirs(DATA_DIR, exist_ok=True)
|
||||
_ensure_default_permissions()
|
||||
_reclassify_existing_tools()
|
||||
p_ensure_default_permissions()
|
||||
p_reclassify_existing_tools()
|
||||
yield
|
||||
|
||||
|
||||
@@ -55,10 +55,10 @@ tools_lib = SubApp("tools", tools_lib_lifespan)
|
||||
# outputs (Gmail, WebFetch); every other built-in is sandboxed by domain.
|
||||
# Must match agent_manager._DEFAULTS so the Settings UI and the agent agree
|
||||
# on what "no policy set" means.
|
||||
_DEFAULT_BUILTIN_POLICIES = {"Bash": "ask"}
|
||||
P_DEFAULT_BUILTIN_POLICIES = {"Bash": "ask"}
|
||||
|
||||
|
||||
def _ensure_default_permissions() -> None:
|
||||
def p_ensure_default_permissions() -> None:
|
||||
"""Seed BUILTIN_PERMISSIONS_PATH so the user's Settings toggles persist
|
||||
cleanly. Without this the file is missing on first run, load returns {},
|
||||
every PUT-from-the-UI overwrites with the partial payload the click
|
||||
@@ -68,15 +68,15 @@ def _ensure_default_permissions() -> None:
|
||||
"""
|
||||
existing = load_builtin_permissions()
|
||||
desired = {
|
||||
t.name: _DEFAULT_BUILTIN_POLICIES.get(t.name, "always_allow")
|
||||
t.name: P_DEFAULT_BUILTIN_POLICIES.get(t.name, "always_allow")
|
||||
for t in BUILTIN_TOOLS
|
||||
}
|
||||
merged = {**desired, **existing}
|
||||
merged = {**desired, **existing}
|
||||
if merged != existing:
|
||||
save_builtin_permissions(merged)
|
||||
p_save_builtin_permissions(merged)
|
||||
|
||||
|
||||
def _reclassify_existing_tools() -> None:
|
||||
def p_reclassify_existing_tools() -> None:
|
||||
"""One-time correction for tools discovered before service rules were integration-scoped: most
|
||||
integrations got mislabeled under a bogus 'Google' group (generic keyword rules applied globally).
|
||||
Recompute services/groups from each tool's stored tool names. Idempotent; rewrites only on change.
|
||||
@@ -87,7 +87,7 @@ def _reclassify_existing_tools() -> None:
|
||||
if not fname.endswith(".json"):
|
||||
continue
|
||||
try:
|
||||
tool = _load(fname[:-5])
|
||||
tool = p_load(fname[:-5])
|
||||
except Exception:
|
||||
continue
|
||||
perms = tool.tool_permissions or {}
|
||||
@@ -96,7 +96,7 @@ def _reclassify_existing_tools() -> None:
|
||||
names = [k for k in perms if not k.startswith("_")]
|
||||
if not names:
|
||||
continue
|
||||
services, service_groups, all_read, all_write = _classify_services(names, tool.name)
|
||||
services, service_groups, all_read, all_write = classify_services(names, tool.name)
|
||||
if perms.get("_services") == services and perms.get("_service_groups") == service_groups:
|
||||
continue
|
||||
perms["_services"] = services
|
||||
@@ -104,7 +104,7 @@ def _reclassify_existing_tools() -> None:
|
||||
perms["_categories"] = {"read": all_read, "write": all_write}
|
||||
tool.tool_permissions = perms
|
||||
try:
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -118,11 +118,11 @@ GOOGLE_USERINFO_URL = "https://www.googleapis.com/oauth2/v2/userinfo"
|
||||
# MCPSearch keystroke; the cache skips re-parsing, revalidated by a per-file stat
|
||||
# signature so any write (ours or external) invalidates instantly. Callers treat
|
||||
# the returned ToolDefinitions as immutable; mutate via _load(tool_id) + _save.
|
||||
_tools_cache: list[ToolDefinition] | None = None
|
||||
_tools_cache_sig: tuple | None = None
|
||||
P_TOOLS_CACHE: list[ToolDefinition] | None = None
|
||||
P_TOOLS_CACHE_SIG: tuple | None = None
|
||||
|
||||
|
||||
def _tools_sig() -> tuple | None:
|
||||
def p_tools_sig() -> tuple | None:
|
||||
if not os.path.exists(DATA_DIR):
|
||||
return ()
|
||||
try:
|
||||
@@ -136,11 +136,12 @@ def _tools_sig() -> tuple | None:
|
||||
return None
|
||||
|
||||
|
||||
def _load_all() -> list[ToolDefinition]:
|
||||
global _tools_cache, _tools_cache_sig
|
||||
sig = _tools_sig()
|
||||
if sig is not None and _tools_cache is not None and sig == _tools_cache_sig:
|
||||
return list(_tools_cache)
|
||||
# Public - called by mcp_pre_flight.py, prompt_context.py, main.py, agent_manager.py
|
||||
def load_all_tools() -> list[ToolDefinition]:
|
||||
global P_TOOLS_CACHE, P_TOOLS_CACHE_SIG
|
||||
sig = p_tools_sig()
|
||||
if sig is not None and P_TOOLS_CACHE is not None and sig == P_TOOLS_CACHE_SIG:
|
||||
return list(P_TOOLS_CACHE)
|
||||
result = []
|
||||
if not os.path.exists(DATA_DIR):
|
||||
return result
|
||||
@@ -149,17 +150,17 @@ def _load_all() -> list[ToolDefinition]:
|
||||
with open(os.path.join(DATA_DIR, fname)) as f:
|
||||
result.append(ToolDefinition(**json.load(f)))
|
||||
if sig is not None:
|
||||
_tools_cache = list(result)
|
||||
_tools_cache_sig = sig
|
||||
P_TOOLS_CACHE = list(result)
|
||||
P_TOOLS_CACHE_SIG = sig
|
||||
return result
|
||||
|
||||
|
||||
def _save(tool: ToolDefinition):
|
||||
def p_save(tool: ToolDefinition):
|
||||
with open(os.path.join(DATA_DIR, f"{tool.id}.json"), "w") as f:
|
||||
json.dump(tool.model_dump(), f, indent=2)
|
||||
|
||||
|
||||
def _load(tool_id: str) -> ToolDefinition:
|
||||
def p_load(tool_id: str) -> ToolDefinition:
|
||||
path = os.path.join(DATA_DIR, f"{tool_id}.json")
|
||||
if not os.path.exists(path):
|
||||
raise HTTPException(status_code=404, detail="Tool not found")
|
||||
@@ -179,7 +180,7 @@ def _load(tool_id: str) -> ToolDefinition:
|
||||
"command": "python",
|
||||
"args": ["-m", "backend.apps.discord_mcp_shim"],
|
||||
}
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return tool
|
||||
|
||||
|
||||
@@ -188,6 +189,7 @@ async def list_builtin_tools():
|
||||
return {"tools": [t.model_dump() for t in BUILTIN_TOOLS]}
|
||||
|
||||
|
||||
# Public - called by agent_manager.py
|
||||
def load_builtin_permissions() -> dict[str, str]:
|
||||
if not os.path.exists(BUILTIN_PERMS_PATH):
|
||||
return {}
|
||||
@@ -195,12 +197,13 @@ def load_builtin_permissions() -> dict[str, str]:
|
||||
return json.load(f)
|
||||
|
||||
|
||||
def save_builtin_permissions(perms: dict[str, str]):
|
||||
def p_save_builtin_permissions(perms: dict[str, str]):
|
||||
os.makedirs(os.path.dirname(BUILTIN_PERMS_PATH), exist_ok=True)
|
||||
with open(BUILTIN_PERMS_PATH, "w") as f:
|
||||
json.dump(perms, f, indent=2)
|
||||
|
||||
|
||||
# Public - called by agent_manager.py
|
||||
def load_trusted_sensitive_paths() -> list[str]:
|
||||
if not os.path.exists(TRUSTED_SENSITIVE_PATHS_PATH):
|
||||
return []
|
||||
@@ -215,6 +218,7 @@ def load_trusted_sensitive_paths() -> list[str]:
|
||||
return [p for p in raw if isinstance(p, str) and p]
|
||||
|
||||
|
||||
# Public - called by agent_manager.py
|
||||
def save_trusted_sensitive_paths(patterns: list[str]):
|
||||
os.makedirs(os.path.dirname(TRUSTED_SENSITIVE_PATHS_PATH), exist_ok=True)
|
||||
seen: list[str] = []
|
||||
@@ -254,14 +258,14 @@ async def update_builtin_permissions(body: dict):
|
||||
for name, policy in body.get("permissions", {}).items():
|
||||
if name in valid_tools and policy in valid_policies:
|
||||
perms[name] = policy
|
||||
save_builtin_permissions(perms)
|
||||
p_save_builtin_permissions(perms)
|
||||
return {"permissions": perms}
|
||||
|
||||
|
||||
@tools_lib.router.get("/list")
|
||||
async def list_tools():
|
||||
tools = []
|
||||
for t in _load_all():
|
||||
for t in load_all_tools():
|
||||
d = t.model_dump()
|
||||
# Heal pre-fix tools whose persisted email is the "{name} account" placeholder so the pill stops reading like a name and falls back to plain "Connected".
|
||||
placeholder = f"{t.name} account"
|
||||
@@ -271,7 +275,7 @@ async def list_tools():
|
||||
return {"tools": tools}
|
||||
|
||||
|
||||
def _connected_html() -> HTMLResponse:
|
||||
def p_connected_html() -> HTMLResponse:
|
||||
"""v1.0.25-style auto-close page. Same markup so the UX is unchanged."""
|
||||
return HTMLResponse("""
|
||||
<html><body>
|
||||
@@ -287,7 +291,7 @@ def _connected_html() -> HTMLResponse:
|
||||
|
||||
@tools_lib.router.get("/{tool_id}")
|
||||
async def get_tool(tool_id: str):
|
||||
return _load(tool_id).model_dump()
|
||||
return p_load(tool_id).model_dump()
|
||||
|
||||
|
||||
@tools_lib.router.post("/create")
|
||||
@@ -301,16 +305,16 @@ async def create_tool(body: ToolCreate):
|
||||
auth_type=body.auth_type,
|
||||
auth_status=body.auth_status,
|
||||
)
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return {"ok": True, "tool": tool.model_dump()}
|
||||
|
||||
|
||||
@tools_lib.router.put("/{tool_id}")
|
||||
async def update_tool(tool_id: str, body: ToolUpdate):
|
||||
tool = _load(tool_id)
|
||||
tool = p_load(tool_id)
|
||||
for k, v in body.model_dump(exclude_none=True).items():
|
||||
setattr(tool, k, v)
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return {"ok": True, "tool": tool.model_dump()}
|
||||
|
||||
|
||||
@@ -324,7 +328,7 @@ async def delete_tool(tool_id: str):
|
||||
|
||||
@tools_lib.router.post("/{tool_id}/discover")
|
||||
async def discover_tools(tool_id: str):
|
||||
tool = _load(tool_id)
|
||||
tool = p_load(tool_id)
|
||||
|
||||
if tool.auth_type == "oauth2" and tool.auth_status == "connected":
|
||||
if tool.oauth_tokens.get("refresh_token"):
|
||||
@@ -353,7 +357,7 @@ async def discover_tools(tool_id: str):
|
||||
command = config.get("command", "")
|
||||
if not command:
|
||||
raise HTTPException(status_code=400, detail="stdio transport requires a 'command' in MCP config")
|
||||
raw_tools = await _discover_mcp_tools_stdio(
|
||||
raw_tools = await discover_mcp_tools_stdio(
|
||||
command=command,
|
||||
args=config.get("args"),
|
||||
env=config.get("env"),
|
||||
@@ -363,13 +367,13 @@ async def discover_tools(tool_id: str):
|
||||
if not url:
|
||||
raise HTTPException(status_code=400, detail="HTTP/SSE transport requires a 'url' in MCP config")
|
||||
if transport == "sse":
|
||||
raw_tools = await _discover_mcp_tools_sse(url, config.get("headers"))
|
||||
raw_tools = await discover_mcp_tools_sse(url, config.get("headers"))
|
||||
else:
|
||||
try:
|
||||
raw_tools = await _discover_mcp_tools_http(url, config.get("headers"))
|
||||
raw_tools = await discover_mcp_tools_http(url, config.get("headers"))
|
||||
except HTTPException:
|
||||
logger.info(f"Streamable HTTP failed for {tool.name}, retrying with SSE transport")
|
||||
raw_tools = await _discover_mcp_tools_sse(url, config.get("headers"))
|
||||
raw_tools = await discover_mcp_tools_sse(url, config.get("headers"))
|
||||
else:
|
||||
raise HTTPException(status_code=400, detail=f"Unsupported MCP transport type: '{transport}'. Use 'stdio', 'http', or 'sse'.")
|
||||
except HTTPException:
|
||||
@@ -382,7 +386,7 @@ async def discover_tools(tool_id: str):
|
||||
raise HTTPException(status_code=502, detail=f"Discovery failed: {msg}")
|
||||
|
||||
tool_names = [t["name"] for t in raw_tools]
|
||||
services, service_groups, all_read, all_write = _classify_services(tool_names, tool.name)
|
||||
services, service_groups, all_read, all_write = classify_services(tool_names, tool.name)
|
||||
permissions: dict[str, Any] = {n: tool.tool_permissions.get(n, "ask") for n in tool_names}
|
||||
permissions["_categories"] = {"read": all_read, "write": all_write}
|
||||
permissions["_services"] = services
|
||||
@@ -391,7 +395,7 @@ async def discover_tools(tool_id: str):
|
||||
permissions["_tool_schemas"] = {t["name"]: t.get("inputSchema") for t in raw_tools if t.get("inputSchema")}
|
||||
|
||||
tool.tool_permissions = permissions
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
|
||||
return {"ok": True, "tool": tool.model_dump()}
|
||||
|
||||
@@ -400,7 +404,7 @@ async def discover_tools(tool_id: str):
|
||||
# Microsoft 365 device-code login (runs in the backend, not the MCP server)
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
_m365_login_processes: dict[str, dict] = {} # tool_id -> {proc, device_code, status, email}
|
||||
P_M365_LOGIN_PROCESSES: dict[str, dict] = {} # tool_id -> {proc, device_code, status, email}
|
||||
|
||||
|
||||
@tools_lib.router.post("/{tool_id}/m365/device-login")
|
||||
@@ -412,8 +416,8 @@ async def m365_device_login(tool_id: str):
|
||||
"""
|
||||
import subprocess
|
||||
|
||||
_load(tool_id)
|
||||
script = _m365_server_script()
|
||||
p_load(tool_id)
|
||||
script = m365_server_script()
|
||||
if not os.path.isfile(script):
|
||||
raise HTTPException(status_code=500, detail="M365 MCP server not installed")
|
||||
|
||||
@@ -427,12 +431,12 @@ async def m365_device_login(tool_id: str):
|
||||
if not cmd:
|
||||
raise HTTPException(status_code=500, detail="No node/electron found")
|
||||
|
||||
env = {**os.environ, **_m365_cache_env()}
|
||||
env = {**os.environ, **m365_cache_env()}
|
||||
if cmd == electron:
|
||||
env["ELECTRON_RUN_AS_NODE"] = "1"
|
||||
|
||||
# Kill any existing login process for this tool
|
||||
existing = _m365_login_processes.pop(tool_id, None)
|
||||
existing = P_M365_LOGIN_PROCESSES.pop(tool_id, None)
|
||||
if existing and existing.get("proc"):
|
||||
try:
|
||||
existing["proc"].kill()
|
||||
@@ -449,7 +453,7 @@ async def m365_device_login(tool_id: str):
|
||||
import threading
|
||||
login_state: dict = {"proc": proc, "status": "waiting_for_code", "device_code": "", "device_code_url": "", "email": None, "output": ""}
|
||||
|
||||
def _read_output():
|
||||
def p_read_output():
|
||||
import re
|
||||
for line in proc.stdout:
|
||||
login_state["output"] += line
|
||||
@@ -469,8 +473,8 @@ async def m365_device_login(tool_id: str):
|
||||
login_state["status"] = "connected"
|
||||
# Try to extract email from output
|
||||
try:
|
||||
import json as _j
|
||||
result = _j.loads(login_state["output"].strip().split("\n")[-1])
|
||||
import json
|
||||
result = json.loads(login_state["output"].strip().split("\n")[-1])
|
||||
if result.get("success"):
|
||||
ud = result.get("userData", {})
|
||||
login_state["email"] = ud.get("userPrincipalName") or ud.get("displayName")
|
||||
@@ -478,20 +482,20 @@ async def m365_device_login(tool_id: str):
|
||||
pass
|
||||
# Update tool status
|
||||
try:
|
||||
t = _load(tool_id)
|
||||
t = p_load(tool_id)
|
||||
t.auth_status = "connected"
|
||||
if login_state.get("email"):
|
||||
t.connected_account_email = login_state["email"]
|
||||
_save(t)
|
||||
p_save(t)
|
||||
except Exception:
|
||||
pass
|
||||
else:
|
||||
login_state["status"] = "error"
|
||||
|
||||
thread = threading.Thread(target=_read_output, daemon=True)
|
||||
thread = threading.Thread(target=p_read_output, daemon=True)
|
||||
thread.start()
|
||||
|
||||
_m365_login_processes[tool_id] = login_state
|
||||
P_M365_LOGIN_PROCESSES[tool_id] = login_state
|
||||
|
||||
# Wait briefly for device code to appear
|
||||
for _ in range(30):
|
||||
@@ -512,13 +516,13 @@ async def m365_device_login(tool_id: str):
|
||||
@tools_lib.router.get("/{tool_id}/m365/device-login/status")
|
||||
async def m365_device_login_status(tool_id: str):
|
||||
"""Poll the status of a pending M365 device-code login."""
|
||||
state = _m365_login_processes.get(tool_id)
|
||||
state = P_M365_LOGIN_PROCESSES.get(tool_id)
|
||||
if not state:
|
||||
# Check if already connected via cached token
|
||||
cache_env = _m365_cache_env()
|
||||
cache_env = m365_cache_env()
|
||||
cache_path = cache_env["MS365_MCP_TOKEN_CACHE_PATH"]
|
||||
if os.path.isfile(cache_path):
|
||||
tool = _load(tool_id)
|
||||
tool = p_load(tool_id)
|
||||
if tool.auth_status == "connected":
|
||||
return {"status": "connected", "email": tool.connected_account_email}
|
||||
return {"status": "no_login_in_progress"}
|
||||
@@ -527,10 +531,10 @@ async def m365_device_login_status(tool_id: str):
|
||||
result: dict = {"status": status}
|
||||
if status == "connected":
|
||||
result["email"] = state.get("email")
|
||||
_m365_login_processes.pop(tool_id, None)
|
||||
P_M365_LOGIN_PROCESSES.pop(tool_id, None)
|
||||
elif status == "error":
|
||||
result["message"] = "Login failed"
|
||||
_m365_login_processes.pop(tool_id, None)
|
||||
P_M365_LOGIN_PROCESSES.pop(tool_id, None)
|
||||
|
||||
return result
|
||||
|
||||
@@ -538,21 +542,21 @@ async def m365_device_login_status(tool_id: str):
|
||||
@tools_lib.router.post("/{tool_id}/m365/disconnect")
|
||||
async def m365_disconnect(tool_id: str):
|
||||
"""Disconnect M365 by clearing the cached token."""
|
||||
tool = _load(tool_id)
|
||||
cache_env = _m365_cache_env()
|
||||
tool = p_load(tool_id)
|
||||
cache_env = m365_cache_env()
|
||||
for path in cache_env.values():
|
||||
if os.path.isfile(path):
|
||||
os.remove(path)
|
||||
tool.auth_status = "configured"
|
||||
tool.connected_account_email = None
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return {"ok": True, "tool": tool.model_dump()}
|
||||
|
||||
|
||||
@tools_lib.router.post("/{tool_id}/oauth/disconnect")
|
||||
async def oauth_disconnect(tool_id: str):
|
||||
"""Clear OAuth tokens and reset auth status so the user can reconnect with a different account."""
|
||||
tool = _load(tool_id)
|
||||
tool = p_load(tool_id)
|
||||
access_token = tool.oauth_tokens.get("access_token")
|
||||
|
||||
if access_token and tool.name.lower() != "notion":
|
||||
@@ -570,7 +574,7 @@ async def oauth_disconnect(tool_id: str):
|
||||
tool.oauth_tokens = {}
|
||||
tool.auth_status = "configured"
|
||||
tool.connected_account_email = None
|
||||
_save(tool)
|
||||
p_save(tool)
|
||||
return {"ok": True, "tool": tool.model_dump()}
|
||||
|
||||
|
||||
@@ -578,8 +582,8 @@ async def oauth_disconnect(tool_id: str):
|
||||
async def oauth_start(tool_id: str):
|
||||
"""Return the OAuth start URL for this tool. All built-in providers
|
||||
proxy through Fly so client_secret values stay server-side."""
|
||||
tool = _load(tool_id)
|
||||
proxied = _proxied_provider_for(tool)
|
||||
tool = p_load(tool_id)
|
||||
proxied = proxied_provider_for(tool)
|
||||
if not proxied:
|
||||
raise HTTPException(
|
||||
status_code=400,
|
||||
@@ -587,11 +591,11 @@ async def oauth_start(tool_id: str):
|
||||
)
|
||||
from backend.config.install_id import get_install_id
|
||||
install_id = get_install_id()
|
||||
_port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
p_port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
params = {
|
||||
"install_id": install_id,
|
||||
"tool_id": tool_id,
|
||||
"local_port": _port,
|
||||
"local_port": p_port,
|
||||
}
|
||||
auth_url = (
|
||||
f"{OPENSWARM_OAUTH_BASE_URL}/api/oauth/{proxied}/start?"
|
||||
@@ -648,7 +652,7 @@ async def oauth_cloud_claim(
|
||||
|
||||
data = resp.json()
|
||||
tokens = data.get("tokens", {}) or {}
|
||||
tool = _load(tool_id)
|
||||
tool = p_load(tool_id)
|
||||
# Google's token endpoint doesn't include the user's email; fetch it
|
||||
# from userinfo so the UI can show "you connected you@gmail.com"
|
||||
# rather than the generic "Google account" placeholder.
|
||||
@@ -663,9 +667,9 @@ async def oauth_cloud_claim(
|
||||
tokens["email"] = info_resp.json().get("email") or ""
|
||||
except Exception as e:
|
||||
logger.warning("Google userinfo lookup post-claim failed: %s", e)
|
||||
_persist_cloud_tokens(tool, tokens)
|
||||
_save(tool)
|
||||
return _connected_html()
|
||||
persist_cloud_tokens(tool, tokens)
|
||||
p_save(tool)
|
||||
return p_connected_html()
|
||||
|
||||
|
||||
@tools_lib.router.post("/google-oauth-token")
|
||||
|
||||
+61
-61
@@ -56,7 +56,7 @@ class FetchBody(BaseModel):
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
def _join_text(parts: list[dict[str, Any]]) -> str:
|
||||
def p_join_text(parts: list[dict[str, Any]]) -> str:
|
||||
out = []
|
||||
for p in parts:
|
||||
if isinstance(p, dict) and p.get("type") == "text":
|
||||
@@ -69,11 +69,11 @@ def _join_text(parts: list[dict[str, Any]]) -> str:
|
||||
# ---------------------------------------------------------------------------
|
||||
|
||||
|
||||
GEMINI_API_BASE = "https://generativelanguage.googleapis.com/v1beta"
|
||||
GEMINI_GROUNDING_MODEL = "gemini-2.5-flash" # cheapest + fastest for grounded calls
|
||||
P_GEMINI_API_BASE = "https://generativelanguage.googleapis.com/v1beta"
|
||||
P_GEMINI_GROUNDING_MODEL = "gemini-2.5-flash" # cheapest + fastest for grounded calls
|
||||
|
||||
OPENAI_API_BASE = "https://api.openai.com/v1"
|
||||
OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_preview
|
||||
P_OPENAI_API_BASE = "https://api.openai.com/v1"
|
||||
P_OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_preview
|
||||
|
||||
# Per-attempt timeouts for the search/fetch cascade. The fast-first ORDERING is
|
||||
# what fixes the ~75s stall (DDG answers in ~1s so the slow grounded backends are
|
||||
@@ -82,16 +82,16 @@ OPENAI_SEARCH_MODEL = "gpt-5-mini" # cheapest model that supports web_search_pr
|
||||
# provider (no response at all) gets cut. Grounded native search legitimately
|
||||
# takes 32-42s (httpx ceiling 45s), so its leash sits at 48s, NOT below 45, or
|
||||
# we'd clip the slow tail of a valid paid call.
|
||||
_DDG_ATTEMPT_TIMEOUT = 6.0 # DDG answers <1s; >6s is a network hang, fall through
|
||||
_GROUNDED_ATTEMPT_TIMEOUT = 48.0 # just above the providers' own 45s httpx timeout
|
||||
P_DDG_ATTEMPT_TIMEOUT = 6.0 # DDG answers <1s; >6s is a network hang, fall through
|
||||
P_GROUNDED_ATTEMPT_TIMEOUT = 48.0 # just above the providers' own 45s httpx timeout
|
||||
# Local httpx + trafilatura fetch of a real page; the fast path for /fetch
|
||||
# (normal pages return in <2s). Set just above WebFetchTool's own 30s httpx
|
||||
# ceiling so a valid-but-slow page still completes locally instead of being
|
||||
# clipped down to a grounded summary; only a truly hung server gets cut.
|
||||
_LOCAL_FETCH_TIMEOUT = 32.0
|
||||
P_LOCAL_FETCH_TIMEOUT = 32.0
|
||||
|
||||
|
||||
async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: bool) -> dict:
|
||||
async def p_gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: bool) -> dict:
|
||||
"""Call Gemini with googleSearch (+ optionally urlContext) grounding.
|
||||
|
||||
Returns {"text": grounded_answer, "chunks": [(title, uri), ...],
|
||||
@@ -106,7 +106,7 @@ async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: b
|
||||
"tools": tools,
|
||||
"generationConfig": {"thinkingConfig": {"thinkingBudget": 0}},
|
||||
}
|
||||
url = f"{GEMINI_API_BASE}/models/{GEMINI_GROUNDING_MODEL}:generateContent"
|
||||
url = f"{P_GEMINI_API_BASE}/models/{P_GEMINI_GROUNDING_MODEL}:generateContent"
|
||||
async with httpx.AsyncClient(timeout=45.0) as client:
|
||||
r = await client.post(
|
||||
url,
|
||||
@@ -133,7 +133,7 @@ async def _gemini_grounded_call(api_key: str, prompt: str, *, use_url_context: b
|
||||
return {"text": text, "chunks": chunks, "queries": queries}
|
||||
|
||||
|
||||
def _format_grounded_as_search_results(grounded: dict, query: str) -> str:
|
||||
def p_format_grounded_as_search_results(grounded: dict, query: str) -> str:
|
||||
"""Format Gemini grounding output to match WebSearchTool's text shape."""
|
||||
lines = []
|
||||
chunks = grounded.get("chunks") or []
|
||||
@@ -147,7 +147,7 @@ def _format_grounded_as_search_results(grounded: dict, query: str) -> str:
|
||||
return "\n\n".join(lines)
|
||||
|
||||
|
||||
def _format_grounded_as_fetch(grounded: dict, url: str) -> str:
|
||||
def p_format_grounded_as_fetch(grounded: dict, url: str) -> str:
|
||||
"""Format Gemini urlContext output to match WebFetchTool's text shape."""
|
||||
parts = [f"Contents of {url}:", ""]
|
||||
text = grounded.get("text") or ""
|
||||
@@ -161,7 +161,7 @@ def _format_grounded_as_fetch(grounded: dict, url: str) -> str:
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
def _resolve_gemini_api_key() -> str | None:
|
||||
def p_resolve_gemini_api_key() -> str | None:
|
||||
"""Pull the AI Studio API key from settings, or None."""
|
||||
try:
|
||||
from backend.apps.settings.store import load_settings
|
||||
@@ -171,7 +171,7 @@ def _resolve_gemini_api_key() -> str | None:
|
||||
return None
|
||||
|
||||
|
||||
def _resolve_openai_api_key() -> str | None:
|
||||
def p_resolve_openai_api_key() -> str | None:
|
||||
try:
|
||||
from backend.apps.settings.store import load_settings
|
||||
s = load_settings()
|
||||
@@ -184,38 +184,38 @@ def _resolve_openai_api_key() -> str | None:
|
||||
# `_refresh_9r_connected()` rather than hit on every search call ,
|
||||
# 9Router's /api/providers is fast but not free, and we already
|
||||
# query it from many places.
|
||||
_NINE_ROUTER_CONNECTED: set[str] = set()
|
||||
_NINE_ROUTER_CACHE_AT: float = 0.0
|
||||
P_NINE_ROUTER_CONNECTED: set[str] = set()
|
||||
P_NINE_ROUTER_CACHE_AT: float = 0.0
|
||||
|
||||
|
||||
async def _refresh_9r_connected() -> set[str]:
|
||||
async def p_refresh_9r_connected() -> set[str]:
|
||||
"""Return the set of currently-active 9Router subscription providers
|
||||
(e.g. {"claude", "codex", "antigravity", "gemini-cli"}). Cached for
|
||||
20s to keep search/fetch endpoints snappy."""
|
||||
global _NINE_ROUTER_CONNECTED, _NINE_ROUTER_CACHE_AT
|
||||
global P_NINE_ROUTER_CONNECTED, P_NINE_ROUTER_CACHE_AT
|
||||
import time
|
||||
now = time.time()
|
||||
if now - _NINE_ROUTER_CACHE_AT < 20.0:
|
||||
return _NINE_ROUTER_CONNECTED
|
||||
if now - P_NINE_ROUTER_CACHE_AT < 20.0:
|
||||
return P_NINE_ROUTER_CONNECTED
|
||||
try:
|
||||
from backend.apps.nine_router.process import is_running, get_providers
|
||||
if not is_running():
|
||||
_NINE_ROUTER_CONNECTED = set()
|
||||
P_NINE_ROUTER_CONNECTED = set()
|
||||
else:
|
||||
conns = await get_providers()
|
||||
_NINE_ROUTER_CONNECTED = {
|
||||
P_NINE_ROUTER_CONNECTED = {
|
||||
c.get("provider")
|
||||
for c in conns
|
||||
if isinstance(c, dict) and c.get("isActive") and c.get("provider")
|
||||
}
|
||||
_NINE_ROUTER_CACHE_AT = now
|
||||
P_NINE_ROUTER_CACHE_AT = now
|
||||
except Exception:
|
||||
# Cache stays; best-effort.
|
||||
pass
|
||||
return _NINE_ROUTER_CONNECTED
|
||||
return P_NINE_ROUTER_CONNECTED
|
||||
|
||||
|
||||
async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> dict:
|
||||
async def p_gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> dict:
|
||||
"""Call 9Router's /v1/messages endpoint with a Gemini model so the
|
||||
user's OAuth subscription (Gemini CLI or Antigravity) covers the
|
||||
search call instead of needing a separate AI Studio API key.
|
||||
@@ -223,12 +223,12 @@ async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> di
|
||||
Routes through Anthropic-shape against 9Router's translator. We
|
||||
request a tool result naturally; the translator surfaces grounded
|
||||
URIs as text + cited sources in the response body. Format-shape
|
||||
matches the existing `_gemini_grounded_call` so downstream
|
||||
matches the existing `p_gemini_grounded_call` so downstream
|
||||
`_format_grounded_as_search_results` works unchanged."""
|
||||
import httpx
|
||||
# Prefer Gemini CLI (broader model coverage). Fall back to
|
||||
# Antigravity if CLI isn't connected.
|
||||
connected = await _refresh_9r_connected()
|
||||
connected = await p_refresh_9r_connected()
|
||||
if "gemini-cli" in connected:
|
||||
model = "gc/gemini-2.5-flash"
|
||||
elif "antigravity" in connected:
|
||||
@@ -269,12 +269,12 @@ async def _gemini_grounded_via_9router(prompt: str, use_url_context: bool) -> di
|
||||
return {"text": text, "chunks": []}
|
||||
|
||||
|
||||
async def _openai_websearch_via_9router(query: str) -> dict:
|
||||
async def p_openai_websearch_via_9router(query: str) -> dict:
|
||||
"""Same idea, but for OpenAI's web_search_preview through Codex's
|
||||
9Router connection. Goes through 9Router's openai-compat endpoint
|
||||
(the responses API) so the user's Codex subscription covers it."""
|
||||
import httpx
|
||||
connected = await _refresh_9r_connected()
|
||||
connected = await p_refresh_9r_connected()
|
||||
if "codex" not in connected:
|
||||
return {}
|
||||
body = {
|
||||
@@ -302,20 +302,20 @@ async def _openai_websearch_via_9router(query: str) -> dict:
|
||||
return {"text": text, "chunks": []}
|
||||
|
||||
|
||||
async def _openai_websearch(api_key: str, query: str) -> dict:
|
||||
async def p_openai_websearch(api_key: str, query: str) -> dict:
|
||||
"""Call OpenAI Responses API with the web_search_preview tool.
|
||||
|
||||
Returns {"text": grounded_answer, "chunks": [(title, uri), ...]}.
|
||||
"""
|
||||
import httpx
|
||||
body = {
|
||||
"model": OPENAI_SEARCH_MODEL,
|
||||
"model": P_OPENAI_SEARCH_MODEL,
|
||||
"input": f"Search the web for: {query}\n\nReturn a concise summary. Cite sources.",
|
||||
"tools": [{"type": "web_search_preview"}],
|
||||
}
|
||||
async with httpx.AsyncClient(timeout=45.0) as client:
|
||||
r = await client.post(
|
||||
f"{OPENAI_API_BASE}/responses",
|
||||
f"{P_OPENAI_API_BASE}/responses",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
)
|
||||
@@ -341,20 +341,20 @@ async def _openai_websearch(api_key: str, query: str) -> dict:
|
||||
return {"text": "".join(text_parts), "chunks": chunks, "queries": [query]}
|
||||
|
||||
|
||||
async def _openai_urlfetch(api_key: str, url: str, prompt: str | None) -> dict:
|
||||
async def p_openai_urlfetch(api_key: str, url: str, prompt: str | None) -> dict:
|
||||
"""Use OpenAI's web_search_preview to fetch/summarize a specific URL."""
|
||||
prompt_text = f"Fetch and summarize the content at: {url}"
|
||||
if prompt:
|
||||
prompt_text += f"\n\nFocus on: {prompt}"
|
||||
import httpx
|
||||
body = {
|
||||
"model": OPENAI_SEARCH_MODEL,
|
||||
"model": P_OPENAI_SEARCH_MODEL,
|
||||
"input": prompt_text,
|
||||
"tools": [{"type": "web_search_preview"}],
|
||||
}
|
||||
async with httpx.AsyncClient(timeout=45.0) as client:
|
||||
r = await client.post(
|
||||
f"{OPENAI_API_BASE}/responses",
|
||||
f"{P_OPENAI_API_BASE}/responses",
|
||||
headers={"Authorization": f"Bearer {api_key}", "Content-Type": "application/json"},
|
||||
json=body,
|
||||
)
|
||||
@@ -391,8 +391,8 @@ async def search(body: SearchBody) -> dict:
|
||||
|
||||
If the primary's own native path fails, we cascade to whichever
|
||||
other provider's credentials are available, then DDG last."""
|
||||
gemini_key = _resolve_gemini_api_key()
|
||||
openai_key = _resolve_openai_api_key()
|
||||
gemini_key = p_resolve_gemini_api_key()
|
||||
openai_key = p_resolve_openai_api_key()
|
||||
primary = (body.primary or "").lower()
|
||||
errors: list[str] = []
|
||||
|
||||
@@ -403,20 +403,20 @@ async def search(body: SearchBody) -> dict:
|
||||
f"Search the web for: {body.query}\n\n"
|
||||
f"Return a concise summary of what you found. Cite sources."
|
||||
)
|
||||
grounded = await _gemini_grounded_call(gemini_key, prompt, use_url_context=False)
|
||||
grounded = await p_gemini_grounded_call(gemini_key, prompt, use_url_context=False)
|
||||
return {
|
||||
"query": body.query,
|
||||
"results": _format_grounded_as_search_results(grounded, body.query),
|
||||
"results": p_format_grounded_as_search_results(grounded, body.query),
|
||||
"backend": "gemini_native",
|
||||
}
|
||||
|
||||
async def try_openai():
|
||||
if not openai_key:
|
||||
return None
|
||||
grounded = await _openai_websearch(openai_key, body.query)
|
||||
grounded = await p_openai_websearch(openai_key, body.query)
|
||||
return {
|
||||
"query": body.query,
|
||||
"results": _format_grounded_as_search_results(grounded, body.query),
|
||||
"results": p_format_grounded_as_search_results(grounded, body.query),
|
||||
"backend": "openai_native",
|
||||
}
|
||||
|
||||
@@ -425,22 +425,22 @@ async def search(body: SearchBody) -> dict:
|
||||
f"Search the web for: {body.query}\n\n"
|
||||
f"Return a concise summary of what you found. Cite sources."
|
||||
)
|
||||
grounded = await _gemini_grounded_via_9router(prompt, use_url_context=False)
|
||||
grounded = await p_gemini_grounded_via_9router(prompt, use_url_context=False)
|
||||
if not grounded.get("text"):
|
||||
return None
|
||||
return {
|
||||
"query": body.query,
|
||||
"results": _format_grounded_as_search_results(grounded, body.query),
|
||||
"results": p_format_grounded_as_search_results(grounded, body.query),
|
||||
"backend": "gemini_subscription",
|
||||
}
|
||||
|
||||
async def try_openai_subscription():
|
||||
grounded = await _openai_websearch_via_9router(body.query)
|
||||
grounded = await p_openai_websearch_via_9router(body.query)
|
||||
if not grounded.get("text"):
|
||||
return None
|
||||
return {
|
||||
"query": body.query,
|
||||
"results": _format_grounded_as_search_results(grounded, body.query),
|
||||
"results": p_format_grounded_as_search_results(grounded, body.query),
|
||||
"backend": "openai_subscription",
|
||||
}
|
||||
|
||||
@@ -473,8 +473,8 @@ async def search(body: SearchBody) -> dict:
|
||||
if primary == "openai":
|
||||
grounded = grounded[2:] + grounded[:2]
|
||||
|
||||
cascade = [("ddg", try_ddg, _DDG_ATTEMPT_TIMEOUT)] + [
|
||||
(name, fn, _GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
|
||||
cascade = [("ddg", try_ddg, P_DDG_ATTEMPT_TIMEOUT)] + [
|
||||
(name, fn, P_GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
|
||||
]
|
||||
|
||||
for name, fn, timeout in cascade:
|
||||
@@ -490,7 +490,7 @@ async def search(body: SearchBody) -> dict:
|
||||
errors.append(f"{name}: {str(e)[:150]}")
|
||||
|
||||
# Everything failed. Be honest about why instead of an empty "no results".
|
||||
connected = await _refresh_9r_connected()
|
||||
connected = await p_refresh_9r_connected()
|
||||
has_subscription = bool(connected & {"codex", "antigravity", "gemini-cli"})
|
||||
if not (gemini_key or openai_key or has_subscription):
|
||||
tail = (
|
||||
@@ -524,8 +524,8 @@ async def fetch(body: FetchBody) -> dict:
|
||||
except SSRFBlocked as exc:
|
||||
from fastapi import HTTPException
|
||||
raise HTTPException(status_code=400, detail=f"Refused: {exc}")
|
||||
gemini_key = _resolve_gemini_api_key()
|
||||
openai_key = _resolve_openai_api_key()
|
||||
gemini_key = p_resolve_gemini_api_key()
|
||||
openai_key = p_resolve_openai_api_key()
|
||||
primary = (body.primary or "").lower()
|
||||
|
||||
async def try_gemini():
|
||||
@@ -534,22 +534,22 @@ async def fetch(body: FetchBody) -> dict:
|
||||
prompt_bits = [f"Fetch and summarize this URL: {body.url}"]
|
||||
if body.prompt:
|
||||
prompt_bits.append(f"Focus on: {body.prompt}")
|
||||
grounded = await _gemini_grounded_call(
|
||||
grounded = await p_gemini_grounded_call(
|
||||
gemini_key, "\n".join(prompt_bits), use_url_context=True,
|
||||
)
|
||||
return {
|
||||
"url": body.url,
|
||||
"content": _format_grounded_as_fetch(grounded, body.url),
|
||||
"content": p_format_grounded_as_fetch(grounded, body.url),
|
||||
"backend": "gemini_native",
|
||||
}
|
||||
|
||||
async def try_openai():
|
||||
if not openai_key:
|
||||
return None
|
||||
grounded = await _openai_urlfetch(openai_key, body.url, body.prompt)
|
||||
grounded = await p_openai_urlfetch(openai_key, body.url, body.prompt)
|
||||
return {
|
||||
"url": body.url,
|
||||
"content": _format_grounded_as_fetch(grounded, body.url),
|
||||
"content": p_format_grounded_as_fetch(grounded, body.url),
|
||||
"backend": "openai_native",
|
||||
}
|
||||
|
||||
@@ -557,14 +557,14 @@ async def fetch(body: FetchBody) -> dict:
|
||||
prompt_bits = [f"Fetch and summarize this URL: {body.url}"]
|
||||
if body.prompt:
|
||||
prompt_bits.append(f"Focus on: {body.prompt}")
|
||||
grounded = await _gemini_grounded_via_9router(
|
||||
grounded = await p_gemini_grounded_via_9router(
|
||||
"\n".join(prompt_bits), use_url_context=True,
|
||||
)
|
||||
if not grounded.get("text"):
|
||||
return None
|
||||
return {
|
||||
"url": body.url,
|
||||
"content": _format_grounded_as_fetch(grounded, body.url),
|
||||
"content": p_format_grounded_as_fetch(grounded, body.url),
|
||||
"backend": "gemini_subscription",
|
||||
}
|
||||
|
||||
@@ -574,12 +574,12 @@ async def fetch(body: FetchBody) -> dict:
|
||||
prompt = f"Fetch this URL and summarize: {body.url}"
|
||||
if body.prompt:
|
||||
prompt += f"\nFocus on: {body.prompt}"
|
||||
grounded = await _openai_websearch_via_9router(prompt)
|
||||
grounded = await p_openai_websearch_via_9router(prompt)
|
||||
if not grounded.get("text"):
|
||||
return None
|
||||
return {
|
||||
"url": body.url,
|
||||
"content": _format_grounded_as_fetch(grounded, body.url),
|
||||
"content": p_format_grounded_as_fetch(grounded, body.url),
|
||||
"backend": "openai_subscription",
|
||||
}
|
||||
|
||||
@@ -597,7 +597,7 @@ async def fetch(body: FetchBody) -> dict:
|
||||
parts = await WebFetchTool().execute(
|
||||
{"url": body.url, "prompt": body.prompt or ""}, None,
|
||||
)
|
||||
text = _join_text(parts)
|
||||
text = p_join_text(parts)
|
||||
local_text = text
|
||||
if text.startswith(("HTTP error", "Error fetching", "Refused to fetch")):
|
||||
return None
|
||||
@@ -615,8 +615,8 @@ async def fetch(body: FetchBody) -> dict:
|
||||
if primary == "openai":
|
||||
grounded = grounded[2:] + grounded[:2]
|
||||
|
||||
cascade = [("local", try_local, _LOCAL_FETCH_TIMEOUT)] + [
|
||||
(name, fn, _GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
|
||||
cascade = [("local", try_local, P_LOCAL_FETCH_TIMEOUT)] + [
|
||||
(name, fn, P_GROUNDED_ATTEMPT_TIMEOUT) for name, fn in grounded
|
||||
]
|
||||
|
||||
errors: list[str] = []
|
||||
|
||||
+55
-49
@@ -10,10 +10,10 @@ from backend.config.paths import AUTH_TOKEN_FILE
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_TOKEN: str = ""
|
||||
TOKEN: str = ""
|
||||
|
||||
|
||||
def _write_atomic(path: str, data: str, mode: int = 0o600) -> None:
|
||||
def p_write_atomic(path: str, data: str, mode: int = 0o600) -> None:
|
||||
"""Atomic write to `path` at the given file mode; never world-readable or half-written."""
|
||||
os.makedirs(os.path.dirname(path), exist_ok=True)
|
||||
tmp = path + ".tmp"
|
||||
@@ -29,100 +29,102 @@ def _write_atomic(path: str, data: str, mode: int = 0o600) -> None:
|
||||
os.replace(tmp, path)
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def init_auth_token() -> str:
|
||||
"""Load the per-install token from disk, or mint one if missing; reused across restarts so Electron's cached copy stays valid."""
|
||||
global _TOKEN
|
||||
global TOKEN
|
||||
try:
|
||||
if os.path.exists(AUTH_TOKEN_FILE):
|
||||
with open(AUTH_TOKEN_FILE, "r", encoding="utf-8") as f:
|
||||
existing = f.read().strip()
|
||||
if existing and 16 <= len(existing) <= 512:
|
||||
_TOKEN = existing
|
||||
TOKEN = existing
|
||||
logger.info(
|
||||
f"auth: reusing existing token from {AUTH_TOKEN_FILE}"
|
||||
)
|
||||
return _TOKEN
|
||||
return TOKEN
|
||||
except Exception as e:
|
||||
logger.warning(f"auth: failed to read existing token, generating new: {e}")
|
||||
|
||||
_TOKEN = secrets.token_urlsafe(32)
|
||||
TOKEN = secrets.token_urlsafe(32)
|
||||
try:
|
||||
_write_atomic(AUTH_TOKEN_FILE, _TOKEN, mode=0o600)
|
||||
p_write_atomic(AUTH_TOKEN_FILE, TOKEN, mode=0o600)
|
||||
logger.info(f"auth: wrote token to {AUTH_TOKEN_FILE} (mode 0600)")
|
||||
except Exception as e:
|
||||
# If we can't write the file, Electron can't read it; log loudly but don't crash.
|
||||
logger.error(f"auth: failed to write token file: {e}")
|
||||
return _TOKEN
|
||||
return TOKEN
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def get_auth_token() -> str:
|
||||
"""Return the current token. Empty string if init_auth_token() hasn't run."""
|
||||
return _TOKEN
|
||||
return TOKEN
|
||||
|
||||
|
||||
class _TokenScrubFilter(logging.Filter):
|
||||
class P_TokenScrubFilter(logging.Filter):
|
||||
"""Logging filter that redacts the install token from log records (defense in depth)."""
|
||||
|
||||
_PLACEHOLDER = "<REDACTED:openswarm-token>"
|
||||
P_PLACEHOLDER = "<REDACTED:openswarm-token>"
|
||||
|
||||
@staticmethod
|
||||
def _args_might_contain_token(args) -> bool:
|
||||
def p_args_might_contain_token(args) -> bool:
|
||||
"""Cheap pre-check; avoids eager %-formatting on the >99% of records that don't mention the token."""
|
||||
if not args:
|
||||
return False
|
||||
items = args if isinstance(args, (tuple, list)) else (args,)
|
||||
for a in items:
|
||||
if isinstance(a, str) and _TOKEN in a:
|
||||
if isinstance(a, str) and TOKEN in a:
|
||||
return True
|
||||
if isinstance(a, dict):
|
||||
for v in a.values():
|
||||
if isinstance(v, str) and _TOKEN in v:
|
||||
if isinstance(v, str) and TOKEN in v:
|
||||
return True
|
||||
return False
|
||||
|
||||
@classmethod
|
||||
def _scrub_args(cls, args):
|
||||
def p_scrub_args(cls, args):
|
||||
"""Scrub token from args while preserving tuple/dict shape; uvicorn's AccessFormatter unpacks args as a 5-tuple and explodes on None."""
|
||||
if args is None:
|
||||
return args
|
||||
if isinstance(args, dict):
|
||||
new_dict = None
|
||||
for k, v in args.items():
|
||||
if isinstance(v, str) and _TOKEN in v:
|
||||
if isinstance(v, str) and TOKEN in v:
|
||||
if new_dict is None:
|
||||
new_dict = dict(args)
|
||||
new_dict[k] = v.replace(_TOKEN, cls._PLACEHOLDER)
|
||||
new_dict[k] = v.replace(TOKEN, cls.P_PLACEHOLDER)
|
||||
return new_dict if new_dict is not None else args
|
||||
if isinstance(args, tuple):
|
||||
new_list = None
|
||||
for i, v in enumerate(args):
|
||||
if isinstance(v, str) and _TOKEN in v:
|
||||
if isinstance(v, str) and TOKEN in v:
|
||||
if new_list is None:
|
||||
new_list = list(args)
|
||||
new_list[i] = v.replace(_TOKEN, cls._PLACEHOLDER)
|
||||
new_list[i] = v.replace(TOKEN, cls.P_PLACEHOLDER)
|
||||
return tuple(new_list) if new_list is not None else args
|
||||
if isinstance(args, str) and _TOKEN in args:
|
||||
return args.replace(_TOKEN, cls._PLACEHOLDER)
|
||||
if isinstance(args, str) and TOKEN in args:
|
||||
return args.replace(TOKEN, cls.P_PLACEHOLDER)
|
||||
return args
|
||||
|
||||
def filter(self, record: logging.LogRecord) -> bool: # pragma: no cover (defensive)
|
||||
if not _TOKEN:
|
||||
if not TOKEN:
|
||||
return True
|
||||
# Fast path: skip eager %-formatting on records that don't mention the token.
|
||||
raw_msg = record.msg if isinstance(record.msg, str) else ""
|
||||
if _TOKEN not in raw_msg and not self._args_might_contain_token(record.args):
|
||||
if TOKEN not in raw_msg and not self.p_args_might_contain_token(record.args):
|
||||
return True
|
||||
# Slow path: in-place args rewrite (preserves shape for AccessFormatter), then re-render to catch tokens buried in custom reprs.
|
||||
try:
|
||||
if isinstance(record.msg, str) and _TOKEN in record.msg:
|
||||
record.msg = record.msg.replace(_TOKEN, self._PLACEHOLDER)
|
||||
if isinstance(record.msg, str) and TOKEN in record.msg:
|
||||
record.msg = record.msg.replace(TOKEN, self.P_PLACEHOLDER)
|
||||
scrubbed = self._scrub_args(record.args)
|
||||
if scrubbed is not record.args:
|
||||
record.args = scrubbed
|
||||
try:
|
||||
rendered = record.getMessage()
|
||||
if _TOKEN in rendered:
|
||||
record.msg = rendered.replace(_TOKEN, self._PLACEHOLDER)
|
||||
if TOKEN in rendered:
|
||||
record.msg = rendered.replace(TOKEN, self.P_PLACEHOLDER)
|
||||
record.args = None
|
||||
except Exception:
|
||||
pass
|
||||
@@ -132,19 +134,20 @@ class _TokenScrubFilter(logging.Filter):
|
||||
return True
|
||||
|
||||
|
||||
_scrubber_installed = False
|
||||
P_SCRUBBER_INSTALLED = False
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def install_token_scrubber() -> None:
|
||||
"""Attach the scrubbing filter to every existing AND future log handler; logger-level filters miss propagated child records."""
|
||||
global _scrubber_installed
|
||||
if _scrubber_installed:
|
||||
global P_SCRUBBER_INSTALLED
|
||||
if P_SCRUBBER_INSTALLED:
|
||||
return
|
||||
|
||||
scrubber = _TokenScrubFilter()
|
||||
scrubber = P_TokenScrubFilter()
|
||||
|
||||
def _attach(handler: logging.Handler) -> None:
|
||||
if not any(isinstance(f, _TokenScrubFilter) for f in handler.filters):
|
||||
def attach(handler: logging.Handler) -> None:
|
||||
if not any(isinstance(f, P_TokenScrubFilter) for f in handler.filters):
|
||||
handler.addFilter(scrubber)
|
||||
|
||||
loggers: list[logging.Logger] = [logging.getLogger()]
|
||||
@@ -153,26 +156,26 @@ def install_token_scrubber() -> None:
|
||||
loggers.append(logger)
|
||||
for logger in loggers:
|
||||
for h in list(logger.handlers):
|
||||
_attach(h)
|
||||
attach(h)
|
||||
|
||||
# Patch addHandler so handlers attached later (uvicorn finishes log config after main.py imports) get the scrubber too.
|
||||
_original_addHandler = logging.Logger.addHandler
|
||||
original_addHandler = logging.Logger.addHandler
|
||||
|
||||
def _patched_addHandler(self: logging.Logger, hdlr: logging.Handler) -> None:
|
||||
_attach(hdlr)
|
||||
return _original_addHandler(self, hdlr)
|
||||
def patched_addHandler(self: logging.Logger, hdlr: logging.Handler) -> None:
|
||||
attach(hdlr)
|
||||
return original_addHandler(self, hdlr)
|
||||
|
||||
logging.Logger.addHandler = _patched_addHandler # type: ignore[assignment]
|
||||
logging.Logger.addHandler = patched_addHandler # type: ignore[assignment]
|
||||
|
||||
root = logging.getLogger()
|
||||
if not any(isinstance(f, _TokenScrubFilter) for f in root.filters):
|
||||
if not any(isinstance(f, P_TokenScrubFilter) for f in root.filters):
|
||||
root.addFilter(scrubber)
|
||||
|
||||
_scrubber_installed = True
|
||||
P_SCRUBBER_INSTALLED = True
|
||||
|
||||
|
||||
# Auth-exempt paths: external redirects with their own nonce/state validation, plus the bootstrap health probe.
|
||||
_AUTH_EXEMPT_EXACT = {
|
||||
P_AUTH_EXEMPT_EXACT = {
|
||||
"/api/subscriptions/callback",
|
||||
"/api/tools/oauth/callback",
|
||||
"/api/tools/oauth/cloud-claim",
|
||||
@@ -191,7 +194,7 @@ _AUTH_EXEMPT_EXACT = {
|
||||
"/api/dev/token",
|
||||
}
|
||||
|
||||
_AUTH_EXEMPT_PREFIX = (
|
||||
P_AUTH_EXEMPT_PREFIX = (
|
||||
# Electron polls /api/health/check before loading the token.
|
||||
"/api/health",
|
||||
# 9Router proxies OpenAI requests with the user's sk-... bearer, not our local token; localhost-only is the gate.
|
||||
@@ -203,11 +206,12 @@ _AUTH_EXEMPT_PREFIX = (
|
||||
)
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def is_path_exempt(path: str) -> bool:
|
||||
"""True if this request path bypasses token auth."""
|
||||
if path in _AUTH_EXEMPT_EXACT:
|
||||
if path in P_AUTH_EXEMPT_EXACT:
|
||||
return True
|
||||
for p in _AUTH_EXEMPT_PREFIX:
|
||||
for p in P_AUTH_EXEMPT_PREFIX:
|
||||
if path.startswith(p):
|
||||
return True
|
||||
return False
|
||||
@@ -224,9 +228,10 @@ def extract_bearer(header_value: str | None) -> str:
|
||||
return ""
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def request_matches_token(request_headers: dict, query_params: dict | None = None) -> bool:
|
||||
"""Validate that an HTTP/WS request carries our token (Bearer, x-openswarm-token, or ?token=); constant-time compare."""
|
||||
if not _TOKEN:
|
||||
if not TOKEN:
|
||||
# Backend not initialized: fail closed. Only test fixtures that bypass main hit this.
|
||||
return False
|
||||
|
||||
@@ -250,13 +255,13 @@ def request_matches_token(request_headers: dict, query_params: dict | None = Non
|
||||
candidates.append(qp_token)
|
||||
|
||||
for candidate in candidates:
|
||||
if secrets.compare_digest(candidate, _TOKEN):
|
||||
if secrets.compare_digest(candidate, TOKEN):
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
# WS Origin allowlist: Electron packaged is file://, dev is localhost:3000, some Electron contexts send bare "null".
|
||||
_ORIGIN_ALLOWLIST_DEV = {
|
||||
P_ORIGIN_ALLOWLIST_DEV = {
|
||||
"http://localhost:3000",
|
||||
"http://127.0.0.1:3000",
|
||||
"file://",
|
||||
@@ -264,12 +269,13 @@ _ORIGIN_ALLOWLIST_DEV = {
|
||||
}
|
||||
|
||||
|
||||
# Public - called by main.py
|
||||
def is_origin_allowed(origin: str | None) -> bool:
|
||||
"""True if the WS connection's Origin header is from our app."""
|
||||
if origin is None:
|
||||
# Native WS client / curl / MCP subprocess: token check still required, so allow.
|
||||
return True
|
||||
if origin in _ORIGIN_ALLOWLIST_DEV:
|
||||
if origin in P_ORIGIN_ALLOWLIST_DEV:
|
||||
return True
|
||||
# Packaged Electron file:// includes paths like file:///Applications/OpenSwarm.app/...; match by prefix.
|
||||
if origin.startswith("file://"):
|
||||
|
||||
@@ -32,8 +32,8 @@ class MainApp:
|
||||
for sub_app in sub_apps:
|
||||
debug(sub_app.name)
|
||||
await stack.enter_async_context(sub_app.lifespan())
|
||||
_port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
print(f"\nCheck out the API docs at: http://127.0.0.1:{_port}/docs\n")
|
||||
port = os.environ.get("OPENSWARM_PORT", "8324")
|
||||
print(f"\nCheck out the API docs at: http://127.0.0.1:{port}/docs\n")
|
||||
yield
|
||||
|
||||
self.app = FastAPI(lifespan=lifespan)
|
||||
|
||||
@@ -7,39 +7,40 @@ import uuid
|
||||
|
||||
from backend.config.paths import DATA_ROOT
|
||||
|
||||
_INSTALL_ID_FILE = os.path.join(DATA_ROOT, "install_id")
|
||||
_cached: str | None = None
|
||||
P_INSTALL_ID_FILE = os.path.join(DATA_ROOT, "install_id")
|
||||
P_CACHED: str | None = None
|
||||
|
||||
|
||||
# Public - called by tools_lib.py, mcp_config.py
|
||||
def get_install_id() -> str:
|
||||
"""Return the persistent install_id, generating and persisting on first call."""
|
||||
global _cached
|
||||
if _cached:
|
||||
return _cached
|
||||
global P_CACHED
|
||||
if P_CACHED:
|
||||
return P_CACHED
|
||||
|
||||
try:
|
||||
with open(_INSTALL_ID_FILE, "r", encoding="utf-8") as f:
|
||||
with open(P_INSTALL_ID_FILE, "r", encoding="utf-8") as f:
|
||||
existing = f.read().strip()
|
||||
if _looks_like_uuid(existing):
|
||||
_cached = existing
|
||||
return _cached
|
||||
if p_looks_like_uuid(existing):
|
||||
P_CACHED = existing
|
||||
return P_CACHED
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
fresh = str(uuid.uuid4())
|
||||
os.makedirs(os.path.dirname(_INSTALL_ID_FILE) or ".", exist_ok=True)
|
||||
fd = os.open(_INSTALL_ID_FILE, os.O_CREAT | os.O_WRONLY | os.O_TRUNC, 0o600)
|
||||
os.makedirs(os.path.dirname(P_INSTALL_ID_FILE) or ".", exist_ok=True)
|
||||
fd = os.open(P_INSTALL_ID_FILE, os.O_CREAT | os.O_WRONLY | os.O_TRUNC, 0o600)
|
||||
try:
|
||||
os.write(fd, fresh.encode("utf-8"))
|
||||
finally:
|
||||
os.close(fd)
|
||||
_cached = fresh
|
||||
return _cached
|
||||
P_CACHED = fresh
|
||||
return P_CACHED
|
||||
|
||||
|
||||
def _looks_like_uuid(s: str) -> bool:
|
||||
def p_looks_like_uuid(s: str) -> bool:
|
||||
if len(s) != 36:
|
||||
return False
|
||||
try:
|
||||
|
||||
@@ -21,12 +21,12 @@ import time
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
_write_lock = threading.Lock()
|
||||
P_WRITE_LOCK = threading.Lock()
|
||||
|
||||
|
||||
def atomic_write_json(path: str, payload, *, indent: int = 2) -> None:
|
||||
directory = os.path.dirname(path) or "."
|
||||
with _write_lock:
|
||||
with P_WRITE_LOCK:
|
||||
os.makedirs(directory, exist_ok=True)
|
||||
fd, tmp = tempfile.mkstemp(prefix=".tmp-", suffix=".json", dir=directory)
|
||||
try:
|
||||
|
||||
@@ -3,20 +3,20 @@
|
||||
import os
|
||||
import sys
|
||||
|
||||
_BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
P_BACKEND_DIR = os.path.dirname(os.path.dirname(os.path.abspath(__file__)))
|
||||
|
||||
_is_packaged = os.environ.get("OPENSWARM_PACKAGED") == "1"
|
||||
P_IS_PACKAGED = os.environ.get("OPENSWARM_PACKAGED") == "1"
|
||||
|
||||
if _is_packaged:
|
||||
if P_IS_PACKAGED:
|
||||
if sys.platform == "darwin":
|
||||
_app_support = os.path.join(os.path.expanduser("~"), "Library", "Application Support", "OpenSwarm")
|
||||
app_support = os.path.join(os.path.expanduser("~"), "Library", "Application Support", "OpenSwarm")
|
||||
elif sys.platform == "win32":
|
||||
_app_support = os.path.join(os.environ.get("APPDATA", os.path.expanduser("~")), "OpenSwarm")
|
||||
app_support = os.path.join(os.environ.get("APPDATA", os.path.expanduser("~")), "OpenSwarm")
|
||||
else:
|
||||
_app_support = os.path.join(os.environ.get("XDG_DATA_HOME", os.path.join(os.path.expanduser("~"), ".local", "share")), "OpenSwarm")
|
||||
DATA_ROOT = os.path.join(_app_support, "data")
|
||||
app_support = os.path.join(os.environ.get("XDG_DATA_HOME", os.path.join(os.path.expanduser("~"), ".local", "share")), "OpenSwarm")
|
||||
DATA_ROOT = os.path.join(app_support, "data")
|
||||
else:
|
||||
DATA_ROOT = os.path.join(_BACKEND_DIR, "data")
|
||||
DATA_ROOT = os.path.join(P_BACKEND_DIR, "data")
|
||||
|
||||
SESSIONS_DIR = os.path.join(DATA_ROOT, "sessions")
|
||||
TOOLS_DIR = os.path.join(DATA_ROOT, "tools")
|
||||
@@ -33,4 +33,4 @@ TRUSTED_SENSITIVE_PATHS_PATH = os.path.join(DATA_ROOT, "trusted_sensitive_paths.
|
||||
# Per-install auth token for the localhost API; see auth.py.
|
||||
AUTH_TOKEN_FILE = os.path.join(DATA_ROOT, "auth.token")
|
||||
|
||||
BACKEND_DIR = _BACKEND_DIR
|
||||
BACKEND_DIR = P_BACKEND_DIR
|
||||
|
||||
+50
-51
@@ -8,13 +8,13 @@ from uuid import uuid4
|
||||
# invisible because nothing configured the 'backend' logger; every debugging
|
||||
# session re-paid that blindness. Idempotent so uvicorn reloads don't stack
|
||||
# handlers; uvicorn's own access logs are untouched.
|
||||
_backend_logger = logging.getLogger("backend")
|
||||
if not _backend_logger.handlers:
|
||||
_h = logging.StreamHandler()
|
||||
_h.setFormatter(logging.Formatter("%(asctime)s %(levelname).1s %(name)s: %(message)s", "%H:%M:%S"))
|
||||
_backend_logger.addHandler(_h)
|
||||
_backend_logger.setLevel(logging.INFO)
|
||||
_backend_logger.propagate = False
|
||||
P_BACKEND_LOGGER = logging.getLogger("backend")
|
||||
if not P_BACKEND_LOGGER.handlers:
|
||||
h = logging.StreamHandler()
|
||||
h.setFormatter(logging.Formatter("%(asctime)s %(levelname).1s %(name)s: %(message)s", "%H:%M:%S"))
|
||||
P_BACKEND_LOGGER.addHandler(h)
|
||||
P_BACKEND_LOGGER.setLevel(logging.INFO)
|
||||
P_BACKEND_LOGGER.propagate = False
|
||||
|
||||
logger = logging.getLogger(__name__)
|
||||
|
||||
@@ -74,12 +74,12 @@ install_token_scrubber()
|
||||
# carries it. Platform-agnostic; wrapped so a settings hiccup never blocks
|
||||
# startup, and the lazy path stays as a fallback.
|
||||
try:
|
||||
import uuid as _uuid
|
||||
from backend.apps.settings.store import load_settings as _load_boot_settings, save_settings as _save_boot_settings
|
||||
_boot_settings = _load_boot_settings()
|
||||
if not getattr(_boot_settings, "installation_id", None):
|
||||
_boot_settings.installation_id = _uuid.uuid4().hex
|
||||
_save_boot_settings(_boot_settings)
|
||||
import uuid
|
||||
from backend.apps.settings.store import load_settings, save_settings
|
||||
boot_settings = load_settings()
|
||||
if not getattr(boot_settings, "installation_id", None):
|
||||
boot_settings.installation_id = uuid.uuid4().hex
|
||||
save_settings(boot_settings)
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
@@ -114,7 +114,7 @@ app.add_middleware(
|
||||
|
||||
|
||||
@app.middleware("http")
|
||||
async def _auth_middleware(request: Request, call_next):
|
||||
async def auth_middleware(request: Request, call_next):
|
||||
"""Reject HTTP requests without our per-install bearer token.
|
||||
|
||||
Exemptions (see `auth.is_path_exempt`):
|
||||
@@ -147,9 +147,9 @@ async def _auth_middleware(request: Request, call_next):
|
||||
# /api/outputs/.../serve/index.html via <iframe src="...">.
|
||||
auth_ok = request_matches_token(headers, query_params=dict(request.query_params))
|
||||
if not auth_ok and x_api_key:
|
||||
import secrets as _s
|
||||
from backend.auth import get_auth_token as _gt
|
||||
auth_ok = _s.compare_digest(x_api_key, _gt() or "\x00")
|
||||
import secrets
|
||||
from backend.auth import get_auth_token
|
||||
auth_ok = secrets.compare_digest(x_api_key, get_auth_token() or "\x00")
|
||||
if not auth_ok:
|
||||
logger.warning(
|
||||
f"auth: rejecting {request.method} {request.url.path} "
|
||||
@@ -184,7 +184,7 @@ async def websocket_session(websocket: WebSocket, session_id: str):
|
||||
things that end a run are: natural completion, explicit
|
||||
`agent:stop`, REST `/close`, or process shutdown.
|
||||
"""
|
||||
if not _ws_auth_ok(websocket):
|
||||
if not p_ws_auth_ok(websocket):
|
||||
return
|
||||
await ws_manager.connect_session(session_id, websocket)
|
||||
try:
|
||||
@@ -257,7 +257,7 @@ async def websocket_session(websocket: WebSocket, session_id: str):
|
||||
# the agent task, that's intentional. See module docstring.
|
||||
ws_manager.disconnect_session(session_id, websocket)
|
||||
|
||||
def _ws_auth_ok(websocket: WebSocket) -> bool:
|
||||
def p_ws_auth_ok(websocket: WebSocket) -> bool:
|
||||
"""Validate token + origin before accepting a WS. Returns True if OK.
|
||||
|
||||
On failure closes with 4401 (custom app-level code) and returns False,
|
||||
@@ -273,8 +273,7 @@ def _ws_auth_ok(websocket: WebSocket) -> bool:
|
||||
logger.warning(f"ws: rejecting connection to {websocket.url.path}, {reason}")
|
||||
# Can't `await websocket.close()` before accept(), so schedule the
|
||||
# close in a task. The client receives a 403 on handshake.
|
||||
import asyncio as _asyncio
|
||||
_asyncio.create_task(websocket.close(code=4401))
|
||||
asyncio.create_task(websocket.close(code=4401))
|
||||
return False
|
||||
return True
|
||||
|
||||
@@ -285,7 +284,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
|
||||
pane. On connect we replay the runtime's ring buffer so a Terminal
|
||||
tab opened mid-session sees the context it missed, then we tail
|
||||
every subsequent line until disconnect."""
|
||||
if not _ws_auth_ok(websocket):
|
||||
if not p_ws_auth_ok(websocket):
|
||||
return
|
||||
await websocket.accept()
|
||||
from backend.apps.outputs.runtime import RUNTIME_MANAGER
|
||||
@@ -318,15 +317,15 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
|
||||
# primed with existing lines before we enter the loop.
|
||||
queue: asyncio.Queue[tuple[str, str]] = asyncio.Queue()
|
||||
|
||||
def _on_line(line) -> None:
|
||||
def on_line(line) -> None:
|
||||
try:
|
||||
queue.put_nowait((line.stream, line.text))
|
||||
except asyncio.QueueFull:
|
||||
pass
|
||||
|
||||
unsubscribe = rt.subscribe(_on_line)
|
||||
unsubscribe = rt.subscribe(on_line)
|
||||
|
||||
def _build_status_frame() -> dict:
|
||||
def build_status_frame() -> dict:
|
||||
return {
|
||||
"event": "runtime:status",
|
||||
"workspace_id": workspace_id,
|
||||
@@ -346,7 +345,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
|
||||
# new-mode preview pointer (Vite dev server); `backend_url` is
|
||||
# the workspace's optional FastAPI backend (old-mode backend.py
|
||||
# OR new-mode post-backend_init.sh).
|
||||
await websocket.send_text(json.dumps(_build_status_frame()))
|
||||
await websocket.send_text(json.dumps(build_status_frame()))
|
||||
while True:
|
||||
stream, text = await queue.get()
|
||||
await websocket.send_text(json.dumps({
|
||||
@@ -361,7 +360,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
|
||||
# to the Vite URL and the preview pane has to know to
|
||||
# switch over. Re-push status after every runtime line.
|
||||
if stream == "runtime":
|
||||
await websocket.send_text(json.dumps(_build_status_frame()))
|
||||
await websocket.send_text(json.dumps(build_status_frame()))
|
||||
except WebSocketDisconnect:
|
||||
pass
|
||||
finally:
|
||||
@@ -370,7 +369,7 @@ async def websocket_runtime_logs(websocket: WebSocket, workspace_id: str):
|
||||
|
||||
@app.websocket("/ws/dashboard")
|
||||
async def websocket_dashboard(websocket: WebSocket):
|
||||
if not _ws_auth_ok(websocket):
|
||||
if not p_ws_auth_ok(websocket):
|
||||
return
|
||||
await ws_manager.connect_global(websocket)
|
||||
try:
|
||||
@@ -446,7 +445,7 @@ async def subscriptions_pending(state: str):
|
||||
}, headers={"Access-Control-Allow-Origin": "*"})
|
||||
|
||||
|
||||
_SUCCESS_HTML = (
|
||||
P_SUCCESS_HTML = (
|
||||
'<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif">'
|
||||
'<div style="text-align:center">'
|
||||
'<div style="width:64px;height:64px;border-radius:50%;background:#22c55e20;display:flex;align-items:center;justify-content:center;margin:0 auto 16px;font-size:32px">✓</div>'
|
||||
@@ -490,7 +489,7 @@ async def subscriptions_callback(request: Request):
|
||||
# Chrome's prefetcher and some extensions speculatively GET URLs.
|
||||
if state and state in COMPLETED_OAUTH:
|
||||
logger.info(f"Duplicate OAuth callback for state {state[:8]}... (already completed)")
|
||||
return HTMLResponse(_SUCCESS_HTML)
|
||||
return HTMLResponse(P_SUCCESS_HTML)
|
||||
logger.warning(f"OAuth callback with unknown state {state[:8] if state else '(empty)'}...")
|
||||
return HTMLResponse('<html><body style="background:#1a1a1a;color:#fff;display:flex;align-items:center;justify-content:center;height:100vh;font-family:sans-serif"><div style="text-align:center"><h2>Session expired</h2><p style="color:#888">Please try connecting again.</p></div></body></html>')
|
||||
|
||||
@@ -508,7 +507,7 @@ async def subscriptions_callback(request: Request):
|
||||
|
||||
mark_oauth_completed(state)
|
||||
logger.info(f"OAuth exchange succeeded for provider={pending.get('provider')}")
|
||||
return HTMLResponse(_SUCCESS_HTML)
|
||||
return HTMLResponse(P_SUCCESS_HTML)
|
||||
|
||||
|
||||
@app.post("/api/browser-agent/run")
|
||||
@@ -550,7 +549,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
valid options instead of activating (anti-hallucination).
|
||||
"""
|
||||
from backend.apps.agents.agent_manager import agent_manager
|
||||
from backend.apps.tools_lib.tools_lib import _load_all as load_all_tools
|
||||
from backend.apps.tools_lib.tools_lib import load_all_tools
|
||||
from backend.apps.tools_lib.mcp_config import sanitize_mcp_server_name
|
||||
|
||||
body = await request.json()
|
||||
@@ -562,7 +561,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
# sanitized server names; values are extra search-hint tokens appended
|
||||
# to the haystack. Only generic synonyms, anything that's already in
|
||||
# the description doesn't need to be listed.
|
||||
_SERVER_SEARCH_ALIASES: dict[str, list[str]] = {
|
||||
P_SERVER_SEARCH_ALIASES: dict[str, list[str]] = {
|
||||
"google-workspace": [
|
||||
"email", "inbox", "mail", "gmail", "calendar", "schedule",
|
||||
"events", "drive", "docs", "sheets", "spreadsheet", "slides",
|
||||
@@ -582,7 +581,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
"youtube": ["video", "transcript", "channel"],
|
||||
}
|
||||
|
||||
def _connected_servers() -> list[dict]:
|
||||
def connected_servers() -> list[dict]:
|
||||
out = []
|
||||
for t in load_all_tools():
|
||||
if not (t.mcp_config and t.enabled and t.auth_status in ("configured", "connected")):
|
||||
@@ -597,7 +596,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
action_names = [str(k) for k in td.keys() if not str(k).startswith("_")]
|
||||
except Exception:
|
||||
pass
|
||||
aliases = _SERVER_SEARCH_ALIASES.get(sanitized, [])
|
||||
aliases = P_SERVER_SEARCH_ALIASES.get(sanitized, [])
|
||||
out.append({
|
||||
"name": sanitized,
|
||||
"description": (t.description or "").strip() or f"{t.name} integration",
|
||||
@@ -606,20 +605,20 @@ async def mcp_meta(action: str, request: Request):
|
||||
})
|
||||
return out
|
||||
|
||||
def _strip_extras(s: dict) -> dict:
|
||||
def strip_extras(s: dict) -> dict:
|
||||
return {k: v for k, v in s.items() if not k.startswith("_")}
|
||||
|
||||
if action == "list":
|
||||
servers = _connected_servers()
|
||||
servers = connected_servers()
|
||||
session = agent_manager.sessions.get(parent_session_id) if parent_session_id else None
|
||||
active_set = set(session.active_mcps) if session else set()
|
||||
active = [{**_strip_extras(s), "status": "active"} for s in servers if s["name"] in active_set]
|
||||
available = [{**_strip_extras(s), "status": "available"} for s in servers if s["name"] not in active_set]
|
||||
active = [{**strip_extras(s), "status": "active"} for s in servers if s["name"] in active_set]
|
||||
available = [{**strip_extras(s), "status": "available"} for s in servers if s["name"] not in active_set]
|
||||
return JSONResponse({"active": active, "available": available})
|
||||
|
||||
if action == "search":
|
||||
query = (body.get("query") or "").strip().lower()
|
||||
servers = _connected_servers()
|
||||
servers = connected_servers()
|
||||
session = agent_manager.sessions.get(parent_session_id) if parent_session_id else None
|
||||
active_set = set(session.active_mcps) if session else set()
|
||||
# Ranking: substring hits across name+description+sub-tool names+
|
||||
@@ -644,7 +643,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
else:
|
||||
score += 1
|
||||
if score:
|
||||
annotated = {**_strip_extras(s), "status": "active" if s["name"] in active_set else "available"}
|
||||
annotated = {**strip_extras(s), "status": "active" if s["name"] in active_set else "available"}
|
||||
scored.append((score, annotated))
|
||||
scored.sort(key=lambda t: (-t[0], 0 if t[1]["status"] == "active" else 1, t[1]["name"]))
|
||||
matches = [s for _, s in scored[:5]]
|
||||
@@ -660,7 +659,7 @@ async def mcp_meta(action: str, request: Request):
|
||||
if not session:
|
||||
return JSONResponse({"error": "session not found"}, status_code=404)
|
||||
|
||||
servers = _connected_servers()
|
||||
servers = connected_servers()
|
||||
valid_names = {s["name"] for s in servers}
|
||||
if server_name not in valid_names:
|
||||
return JSONResponse({"status": "unknown_server", "available": sorted(valid_names)})
|
||||
@@ -681,8 +680,8 @@ async def mcp_meta(action: str, request: Request):
|
||||
if session.sdk_session_id:
|
||||
session.needs_fresh_session = True
|
||||
try:
|
||||
from backend.apps.agents.core.ws_manager import ws_manager as _ws
|
||||
await _ws.send_to_session(parent_session_id, "agent:status", {
|
||||
from backend.apps.agents.core.ws_manager import ws_manager
|
||||
await ws_manager.send_to_session(parent_session_id, "agent:status", {
|
||||
"session_id": parent_session_id,
|
||||
"status": session.status,
|
||||
"session": session.model_dump(mode="json"),
|
||||
@@ -749,14 +748,14 @@ async def session_compact(session_id: str):
|
||||
only sets the marker; the button is the user opting into the cost).
|
||||
"""
|
||||
from backend.apps.agents.agent_manager import agent_manager
|
||||
from backend.apps.agents.core.ws_manager import ws_manager as _ws
|
||||
from backend.apps.agents.core.ws_manager import ws_manager
|
||||
session = agent_manager.sessions.get(session_id)
|
||||
if not session:
|
||||
return JSONResponse({"error": "session not found"}, status_code=404)
|
||||
did_compact = agent_manager._maybe_compact(session, force=True)
|
||||
if did_compact:
|
||||
session.needs_fresh_session = True
|
||||
await _ws.send_to_session(session_id, "agent:context_status", {
|
||||
await ws_manager.send_to_session(session_id, "agent:context_status", {
|
||||
"session_id": session_id,
|
||||
"reason": "compacted_manual" if did_compact else "noop",
|
||||
"compacted_through_msg_id": session.compacted_through_msg_id,
|
||||
@@ -768,7 +767,7 @@ async def session_compact(session_id: str):
|
||||
async def session_clear(session_id: str):
|
||||
"""Wipe the session's UI history AND its SDK convo state (/clear slash cmd, Reset history button)."""
|
||||
from backend.apps.agents.agent_manager import agent_manager
|
||||
from backend.apps.agents.core.ws_manager import ws_manager as _ws
|
||||
from backend.apps.agents.core.ws_manager import ws_manager
|
||||
from backend.apps.agents.core.models import MessageBranch
|
||||
session = agent_manager.sessions.get(session_id)
|
||||
if not session:
|
||||
@@ -784,12 +783,12 @@ async def session_clear(session_id: str):
|
||||
session.branches = {"main": MessageBranch(id="main")}
|
||||
session.active_branch_id = "main"
|
||||
session.tool_group_meta = {}
|
||||
await _ws.send_to_session(session_id, "agent:status", {
|
||||
await ws_manager.send_to_session(session_id, "agent:status", {
|
||||
"session_id": session_id,
|
||||
"status": session.status,
|
||||
"session": session.model_dump(mode="json"),
|
||||
})
|
||||
await _ws.send_to_session(session_id, "agent:context_status", {
|
||||
await ws_manager.send_to_session(session_id, "agent:context_status", {
|
||||
"session_id": session_id,
|
||||
"reason": "cleared",
|
||||
})
|
||||
@@ -841,7 +840,7 @@ if __name__ == "__main__":
|
||||
|
||||
import uvicorn.config
|
||||
|
||||
class _ReadyServer(uvicorn.Server):
|
||||
class P_ReadyServer(uvicorn.Server):
|
||||
"""Subclass that prints a machine-readable READY line on startup."""
|
||||
async def startup(self, sockets=None):
|
||||
await super().startup(sockets)
|
||||
@@ -851,6 +850,6 @@ if __name__ == "__main__":
|
||||
uvicorn.run("backend.main:app", host=args.host, port=args.port, reload=True)
|
||||
else:
|
||||
config = uvicorn.Config("backend.main:app", host=args.host, port=args.port)
|
||||
server = _ReadyServer(config)
|
||||
server = P_ReadyServer(config)
|
||||
import asyncio
|
||||
asyncio.run(server.serve())
|
||||
|
||||
@@ -12,7 +12,7 @@ import pytest
|
||||
|
||||
|
||||
@pytest.fixture(autouse=True)
|
||||
def _isolate_browser_state(monkeypatch):
|
||||
def isolate_browser_state(monkeypatch):
|
||||
skills_dir = tempfile.mkdtemp(prefix="os_skills_")
|
||||
metrics_dir = tempfile.mkdtemp(prefix="os_metrics_")
|
||||
playbook_dir = tempfile.mkdtemp(prefix="os_playbook_")
|
||||
@@ -20,7 +20,7 @@ def _isolate_browser_state(monkeypatch):
|
||||
monkeypatch.setenv("OPENSWARM_BROWSER_METRICS_DIR", metrics_dir)
|
||||
monkeypatch.setenv("OPENSWARM_BROWSER_PLAYBOOK_DIR", playbook_dir)
|
||||
|
||||
def _reset():
|
||||
def reset():
|
||||
for mod in ("browser_skills", "browser_playbook"):
|
||||
try:
|
||||
m = __import__(f"backend.apps.agents.browser.{mod}", fromlist=[mod])
|
||||
@@ -30,10 +30,10 @@ def _isolate_browser_state(monkeypatch):
|
||||
# metrics caches its dir at first use; drop it so each test writes
|
||||
# where ITS env var points, not where the first test's pointed
|
||||
try:
|
||||
from backend.apps.agents.browser import browser_metrics as _bm
|
||||
_bm._metrics_dir_cache = None
|
||||
from backend.apps.agents.browser import browser_metrics
|
||||
browser_metrics.P_METRICS_DIR_CACHE = None
|
||||
except Exception:
|
||||
pass
|
||||
_reset()
|
||||
reset()
|
||||
yield
|
||||
_reset()
|
||||
reset()
|
||||
|
||||
@@ -19,13 +19,13 @@ def client():
|
||||
"""Returns a TestClient pre-loaded with the local backend's auth token
|
||||
so the LocalAuthMiddleware doesn't reject our requests with 401."""
|
||||
import backend.auth as auth_mod
|
||||
if not auth_mod._TOKEN:
|
||||
if not auth_mod.TOKEN:
|
||||
# Tests sometimes run without backend.main's startup hook firing.
|
||||
# Generate a token directly so request_matches_token has something
|
||||
# to compare against.
|
||||
import secrets
|
||||
auth_mod._TOKEN = secrets.token_urlsafe(32)
|
||||
return TestClient(app, headers={"Authorization": f"Bearer {auth_mod._TOKEN}"})
|
||||
auth_mod.TOKEN = secrets.token_urlsafe(32)
|
||||
return TestClient(app, headers={"Authorization": f"Bearer {auth_mod.TOKEN}"})
|
||||
|
||||
|
||||
@pytest.fixture
|
||||
@@ -206,7 +206,7 @@ def test_dev_token_is_dev_only():
|
||||
os.environ.pop("OPENSWARM_PACKAGED", None)
|
||||
r = noauth.get("/api/dev/token")
|
||||
assert r.status_code == 200
|
||||
assert r.json()["token"] == auth_mod._TOKEN
|
||||
assert r.json()["token"] == auth_mod.TOKEN
|
||||
|
||||
os.environ["OPENSWARM_PACKAGED"] = "1"
|
||||
try:
|
||||
|
||||
Reference in New Issue
Block a user