mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-05 00:13:01 +00:00
Member-mode resolution ranked accounts by last_used_at first, so a connection added by Connect to continue (never used) lost to any older used account. Accounts now rank by the later of last use and creation.
697 lines
34 KiB
Python
697 lines
34 KiB
Python
"""Runtime use of connections: tool execution, sharing modes and scheduled sync."""
|
|
|
|
from __future__ import annotations
|
|
|
|
import json
|
|
import logging
|
|
from contextlib import contextmanager
|
|
from types import SimpleNamespace
|
|
from unittest.mock import MagicMock, patch
|
|
|
|
import pytest
|
|
from sqlalchemy import text
|
|
|
|
from docsgpt.security.encryption import encrypt_json
|
|
|
|
|
|
@contextmanager
|
|
def _service_db(conn):
|
|
@contextmanager
|
|
def _yield():
|
|
yield conn
|
|
|
|
with patch.multiple("docsgpt.connectors.service", db_session=_yield, db_readonly=_yield), \
|
|
patch.multiple("docsgpt.connectors.resolve", db_readonly=_yield):
|
|
yield
|
|
|
|
|
|
def _connection(conn, user="alice", provider="telegram", status="connected", secrets=None,
|
|
auth_kind="api_key", server_url=None) -> str:
|
|
return str(conn.execute(
|
|
text(
|
|
"INSERT INTO connector_sessions (user_id, provider, connector_key, auth_kind, status, "
|
|
"account_label, server_url, encrypted_credentials) VALUES (:u, :p, :p, :a, :s, :l, :url, :e) RETURNING id"
|
|
),
|
|
{"u": user, "p": provider, "a": auth_kind, "s": status, "l": f"{user}-label", "url": server_url,
|
|
"e": encrypt_json(secrets or {"credentials": {"token": f"{user}-token"}}, user)},
|
|
).scalar())
|
|
|
|
|
|
def _tool(connection_id, *, user="alice", name="telegram", mode="owner", tool_id="tool-1"):
|
|
return {
|
|
"id": tool_id,
|
|
"user_id": user,
|
|
"name": name,
|
|
"config": {},
|
|
"actions": [{"name": "telegram_send_message", "active": True, "require_approval": False}],
|
|
"connection_id": connection_id,
|
|
"credential_mode": mode,
|
|
}
|
|
|
|
|
|
def _call(action="telegram_send_message"):
|
|
return SimpleNamespace(id="call-1", name=action, arguments="{}", thought_signature=None)
|
|
|
|
|
|
def _executor(user="alice", headless=False):
|
|
from docsgpt.agents.tool_executor import ToolExecutor
|
|
|
|
executor = ToolExecutor(user=user, headless=headless)
|
|
return executor
|
|
|
|
|
|
def _pause(executor, tool):
|
|
with patch("docsgpt.agents.tool_executor.ToolActionParser") as parser:
|
|
parser.return_value.parse_args.return_value = ("t1", "telegram_send_message", {})
|
|
return executor.check_pause({"t1": tool}, _call(), "OpenAILLM")
|
|
|
|
|
|
class TestResolution:
|
|
def test_owner_mode_uses_owner_connection(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(cid), "bob")
|
|
assert resolved.available and resolved.connection_id == cid
|
|
assert resolved.delegated is True
|
|
assert resolved.connector_name == "Telegram"
|
|
|
|
def test_member_mode_uses_invokers_connection(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
owner = _connection(pg_conn)
|
|
bobs = _connection(pg_conn, user="bob")
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(owner, mode="member"), "bob")
|
|
assert resolved.connection_id == bobs
|
|
assert resolved.delegated is False
|
|
|
|
def test_member_with_several_accounts_uses_the_one_last_used(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
owner = _connection(pg_conn)
|
|
older = _connection(pg_conn, user="bob")
|
|
pg_conn.execute(text(
|
|
"UPDATE connector_sessions SET account_label = 'bob-home', last_used_at = now() - interval '1 day' "
|
|
"WHERE id = CAST(:i AS uuid)"
|
|
), {"i": older})
|
|
newer = _connection(pg_conn, user="bob")
|
|
pg_conn.execute(text(
|
|
"UPDATE connector_sessions SET last_used_at = now() WHERE id = CAST(:i AS uuid)"
|
|
), {"i": newer})
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(owner, mode="member"), "bob")
|
|
assert resolved.available and resolved.connection_id == newer
|
|
|
|
def test_member_gets_the_account_they_just_connected(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
owner = _connection(pg_conn)
|
|
used = _connection(pg_conn, user="bob")
|
|
pg_conn.execute(text(
|
|
"UPDATE connector_sessions SET account_label = 'bob-work', created_at = now() - interval '30 days', "
|
|
"updated_at = now() - interval '30 days', last_used_at = now() - interval '1 hour' "
|
|
"WHERE id = CAST(:i AS uuid)"
|
|
), {"i": used})
|
|
# Added by "Connect to continue": never used yet.
|
|
added = _connection(pg_conn, user="bob")
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(owner, mode="member"), "bob")
|
|
assert resolved.available and resolved.connection_id == added
|
|
|
|
def test_member_mode_without_own_connection_is_unavailable(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
owner = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(owner, mode="member"), "bob")
|
|
assert resolved.available is False
|
|
assert resolved.connector_name == "Telegram"
|
|
|
|
def test_member_mode_matches_mcp_server(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
owner = _connection(pg_conn, provider="mcp:https://a.example.com", auth_kind="mcp_oauth",
|
|
server_url="https://a.example.com", secrets={"tokens": {"access_token": "x"}})
|
|
_connection(pg_conn, user="bob", provider="mcp:https://b.example.com", auth_kind="mcp_oauth",
|
|
server_url="https://b.example.com", secrets={"tokens": {"access_token": "y"}})
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(owner, name="mcp_tool", mode="member"), "bob")
|
|
assert resolved.available is False
|
|
|
|
def test_resource_cannot_borrow_another_users_connection(self, pg_conn):
|
|
from docsgpt.connectors.resolve import resolve_connection
|
|
|
|
mallorys_target = _connection(pg_conn, user="victim")
|
|
with _service_db(pg_conn):
|
|
resolved = resolve_connection(_tool(mallorys_target, user="mallory"), "mallory")
|
|
assert resolved.available is False and resolved.row is None
|
|
|
|
|
|
class TestExecutor:
|
|
def test_credentials_come_from_the_connection(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
executor = _executor()
|
|
tool = _tool(cid)
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
executor._get_or_load_tool(tool, "t1", "telegram_send_message")
|
|
config = manager.return_value.load_tool.call_args.kwargs["tool_config"]
|
|
assert config["token"] == "alice-token"
|
|
assert "encrypted_credentials" not in config
|
|
|
|
def test_owner_mode_delegation_is_audited(self, pg_conn, caplog):
|
|
cid = _connection(pg_conn)
|
|
executor = _executor(user="bob")
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager"), \
|
|
caplog.at_level(logging.INFO, logger="docsgpt.connectors.resolve"):
|
|
executor._get_or_load_tool(_tool(cid), "t1", "telegram_send_message")
|
|
record = next(r for r in caplog.records if r.message == "tool_credential_delegation")
|
|
assert record.connection_id == cid and record.invoker == "bob"
|
|
|
|
def test_needs_reconnect_pauses_on_connect_card(self, pg_conn):
|
|
cid = _connection(pg_conn, status="reconnect_needed")
|
|
with _service_db(pg_conn):
|
|
pause = _pause(_executor(), _tool(cid))
|
|
assert pause["pause_type"] == "awaiting_approval"
|
|
# The caller's own connection: the card can reconnect it in place.
|
|
assert pause["connection_required"] == {
|
|
"connector_key": "telegram", "connector_name": "Telegram", "status": "reconnect_needed",
|
|
"connection_id": cid, "owner_account": False,
|
|
}
|
|
|
|
def test_owners_broken_account_is_not_handed_to_the_member(self, pg_conn):
|
|
cid = _connection(pg_conn, status="reconnect_needed")
|
|
with _service_db(pg_conn):
|
|
pause = _pause(_executor(user="bob"), _tool(cid))
|
|
required = pause["connection_required"]
|
|
assert required["owner_account"] is True
|
|
assert "connection_id" not in required
|
|
|
|
def test_member_without_connection_pauses(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
pause = _pause(_executor(user="bob"), _tool(cid, mode="member"))
|
|
assert pause["connection_required"]["status"] == "missing"
|
|
|
|
def test_headless_run_is_denied_not_paused(self, pg_conn):
|
|
cid = _connection(pg_conn, status="disconnected")
|
|
with _service_db(pg_conn):
|
|
pause = _pause(_executor(headless=True), _tool(cid))
|
|
assert pause["pause_type"] == "headless_denied"
|
|
assert pause["error_type"] == "connection_required"
|
|
|
|
def test_connected_tool_does_not_pause(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
assert _pause(_executor(), _tool(cid)) is None
|
|
|
|
def test_mcp_tool_gets_connection_id_not_tokens(self, pg_conn):
|
|
cid = _connection(pg_conn, provider="mcp:https://m.example.com", auth_kind="mcp_oauth",
|
|
server_url="https://m.example.com", secrets={"tokens": {"access_token": "secret"}})
|
|
tool = {**_tool(cid, name="mcp_tool"), "config": {"server_url": "https://m.example.com/mcp",
|
|
"auth_type": "oauth"}}
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
_executor()._get_or_load_tool(tool, "t1", "search")
|
|
config = manager.return_value.load_tool.call_args.kwargs["tool_config"]
|
|
assert config["connection_id"] == cid
|
|
assert "secret" not in str(config)
|
|
|
|
def test_stored_connection_id_in_config_is_ignored(self, pg_conn):
|
|
"""Only a resolved connection reaches the tool; a config value never does."""
|
|
victim = _connection(pg_conn, user="victim", provider="mcp:https://m.example.com", auth_kind="mcp_oauth",
|
|
server_url="https://m.example.com", secrets={"tokens": {"access_token": "v"}})
|
|
tool = {**_tool(None, name="mcp_tool"), "config": {"server_url": "https://m.example.com/mcp",
|
|
"auth_type": "oauth", "connection_id": victim}}
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
_executor()._get_or_load_tool(tool, "t1", "search")
|
|
config = manager.return_value.load_tool.call_args.kwargs["tool_config"]
|
|
assert "connection_id" not in config
|
|
|
|
|
|
def _telegram_tool(connection_id, *, user="alice", mode="owner"):
|
|
from docsgpt.agents.tools.telegram import TelegramTool
|
|
from docsgpt.connectors.service import _transform_actions
|
|
|
|
return {
|
|
**_tool(connection_id, user=user, mode=mode),
|
|
"actions": _transform_actions(TelegramTool({}).get_actions_metadata()),
|
|
}
|
|
|
|
|
|
def _run_send(executor, tool, arguments):
|
|
with patch("docsgpt.agents.tool_executor.ToolActionParser") as parser, \
|
|
patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
parser.return_value.parse_args.return_value = ("t1", "telegram_send_message", arguments)
|
|
gen = executor.execute({"t1": tool}, _call(), "OpenAILLM")
|
|
while True:
|
|
try:
|
|
next(gen)
|
|
except StopIteration:
|
|
break
|
|
return manager.return_value.load_tool.return_value.execute_action.call_args
|
|
|
|
|
|
class TestTelegramDefaultChat:
|
|
def test_the_connection_offers_a_default_chat_field(self):
|
|
from docsgpt.connectors import catalog
|
|
|
|
fields = {f.key: f for f in catalog.get_definition("telegram").credential_fields}
|
|
chat = fields["chat_id"]
|
|
assert chat.secret is False and chat.required is False
|
|
assert chat.parameter == "chat_id"
|
|
assert chat.hint
|
|
assert chat.to_dict()["hint"] == chat.hint
|
|
|
|
def test_the_model_is_not_asked_for_a_chat_the_connection_sets(self, pg_conn):
|
|
cid = _connection(pg_conn, secrets={"credentials": {"token": "t", "chat_id": "-1001"}})
|
|
with _service_db(pg_conn):
|
|
functions = _executor().prepare_tools_for_llm({"t1": _telegram_tool(cid)})
|
|
by_name = {f["function"]["name"]: f["function"]["parameters"] for f in functions}
|
|
assert "chat_id" not in by_name["telegram_send_message"]["properties"]
|
|
assert "chat_id" not in by_name["telegram_send_image"]["properties"]
|
|
|
|
def test_without_a_default_chat_the_model_still_names_one(self, pg_conn):
|
|
cid = _connection(pg_conn, secrets={"credentials": {"token": "t"}})
|
|
with _service_db(pg_conn):
|
|
functions = _executor().prepare_tools_for_llm({"t1": _telegram_tool(cid)})
|
|
params = {f["function"]["name"]: f["function"]["parameters"] for f in functions}
|
|
assert "chat_id" in params["telegram_send_message"]["properties"]
|
|
|
|
def test_the_default_chat_wins_over_what_the_model_sends(self, pg_conn):
|
|
cid = _connection(pg_conn, secrets={"credentials": {"token": "t", "chat_id": "-1001"}})
|
|
with _service_db(pg_conn):
|
|
call = _run_send(_executor(), _telegram_tool(cid), {"text": "hi", "chat_id": "666"})
|
|
assert call.kwargs == {"text": "hi", "chat_id": "-1001"}
|
|
|
|
def test_each_member_uses_their_own_chat(self, pg_conn):
|
|
owner = _connection(pg_conn, secrets={"credentials": {"token": "t", "chat_id": "-1001"}})
|
|
_connection(pg_conn, user="bob", secrets={"credentials": {"token": "b", "chat_id": "-2002"}})
|
|
with _service_db(pg_conn):
|
|
call = _run_send(_executor(user="bob"), _telegram_tool(owner, mode="member"), {"text": "hi"})
|
|
assert call.kwargs == {"text": "hi", "chat_id": "-2002"}
|
|
|
|
|
|
class TestAccountsTellApartForTheModel:
|
|
@staticmethod
|
|
def _two_bots(pg_conn, names):
|
|
tools = {}
|
|
for index, name in enumerate(names):
|
|
cid = _connection(pg_conn)
|
|
pg_conn.execute(text(
|
|
"UPDATE connector_sessions SET account_label = :l, account_name = :n WHERE id = CAST(:i AS uuid)"
|
|
), {"l": f"…{index}abc", "n": name, "i": cid})
|
|
tools[f"t{index}"] = {**_telegram_tool(cid), "id": f"tool-{index}"}
|
|
return tools
|
|
|
|
def test_named_accounts_name_the_functions(self, pg_conn):
|
|
tools = self._two_bots(pg_conn, ["Alerts bot", "Ops: on-call!"])
|
|
with _service_db(pg_conn):
|
|
executor = _executor()
|
|
functions = {f["function"]["name"]: f["function"] for f in executor.prepare_tools_for_llm(tools)}
|
|
assert {"telegram_send_message_alerts_bot", "telegram_send_message_ops_on_call"} <= set(functions)
|
|
assert "Alerts bot" in functions["telegram_send_message_alerts_bot"]["description"]
|
|
assert executor._name_to_tool["telegram_send_message_ops_on_call"] == ("t1", "telegram_send_message")
|
|
|
|
def test_unnamed_accounts_use_their_labels(self, pg_conn):
|
|
tools = self._two_bots(pg_conn, [None, None])
|
|
with _service_db(pg_conn):
|
|
names = {f["function"]["name"] for f in _executor().prepare_tools_for_llm(tools)}
|
|
assert {"telegram_send_message_0abc", "telegram_send_message_1abc"} <= names
|
|
|
|
def test_different_services_are_named_after_the_service(self, pg_conn):
|
|
tools = {}
|
|
for index, (host, name) in enumerate((("a.example.com", "Wiki"), ("b.example.com", "Tracker"))):
|
|
cid = _connection(pg_conn, provider=f"mcp:https://{host}", auth_kind="mcp_oauth",
|
|
server_url=f"https://{host}", secrets={"tokens": {"access_token": "x"}})
|
|
pg_conn.execute(text("UPDATE connector_sessions SET connector_key = 'custom_mcp', display_name = :n "
|
|
"WHERE id = CAST(:i AS uuid)"), {"n": name, "i": cid})
|
|
tools[f"t{index}"] = {**_tool(cid, name="mcp_tool", tool_id=f"tool-{index}"),
|
|
"actions": [{"name": "search", "description": "Search", "active": True}]}
|
|
with _service_db(pg_conn):
|
|
functions = {f["function"]["name"]: f["function"] for f in _executor().prepare_tools_for_llm(tools)}
|
|
assert set(functions) == {"search_wiki", "search_tracker"}
|
|
assert functions["search_wiki"]["description"] == "Search (Wiki)"
|
|
|
|
def test_names_stay_within_provider_limits(self, pg_conn):
|
|
tools = self._two_bots(pg_conn, ["x" * 80, "x" * 80])
|
|
with _service_db(pg_conn):
|
|
names = [f["function"]["name"] for f in _executor().prepare_tools_for_llm(tools)]
|
|
assert len(names) == len(set(names))
|
|
assert all(len(n) <= 64 and n.replace("_", "").replace("-", "").isalnum() for n in names)
|
|
|
|
|
|
class TestScheduledSync:
|
|
def test_connector_sources_with_a_connection_are_dispatched(self, pg_conn):
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="google_drive", auth_kind="oauth",
|
|
secrets={"token_info": {"access_token": "a"}})
|
|
pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency, connection_id, remote_data) "
|
|
"VALUES ('alice', 'Drive', 'connector:file', 'weekly', CAST(:c AS uuid), '{\"provider\": \"google_drive\"}')"
|
|
), {"c": cid})
|
|
pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency) VALUES ('alice', 'Old', 'connector:file', 'weekly')"
|
|
))
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), patch(
|
|
"docsgpt.api.user.tasks.sync_connector_source.delay"
|
|
) as delay:
|
|
counts = worker.sync_worker(MagicMock(), "weekly")
|
|
assert delay.call_count == 1
|
|
assert counts["sync_dispatched"] == 1
|
|
assert counts["sync_skipped"] == 1
|
|
|
|
def test_paused_repository_is_skipped_until_reconnected(self, pg_conn):
|
|
"""A GitHub or S3 source paused for reconnect is not retried (and failed) on every schedule."""
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="github", status="reconnect_needed",
|
|
secrets={"credentials": {"access_token": "revoked"}})
|
|
pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency, connection_id, remote_data, metadata) "
|
|
"VALUES ('alice', 'acme/api', 'github', 'daily', CAST(:c AS uuid), '{\"repo_url\": \"acme/api\"}', "
|
|
"'{\"sync_state\": \"paused_reconnect\"}')"
|
|
), {"c": cid})
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), patch.object(worker, "sync") as sync:
|
|
counts = worker.sync_worker(MagicMock(), "daily")
|
|
sync.assert_not_called()
|
|
assert counts["sync_skipped"] == 1
|
|
|
|
def test_paused_connection_is_not_synced(self, pg_conn):
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="google_drive", auth_kind="oauth", status="reconnect_needed",
|
|
secrets={"token_info": {}})
|
|
source = pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency, connection_id) "
|
|
"VALUES ('alice', 'Drive', 'connector:file', 'weekly', CAST(:c AS uuid)) RETURNING id"
|
|
), {"c": cid}).scalar()
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), patch.object(worker, "ingest_connector") as ingest:
|
|
result = worker.sync_connector_source(MagicMock(), str(source))
|
|
assert result == {"status": "paused"}
|
|
ingest.assert_not_called()
|
|
|
|
def test_disabled_connector_is_not_synced(self, pg_conn):
|
|
from docsgpt import worker
|
|
from docsgpt.storage.db.repositories.connector_policies import ConnectorPoliciesRepository
|
|
|
|
cid = _connection(pg_conn, provider="google_drive", auth_kind="oauth",
|
|
secrets={"token_info": {"access_token": "a"}})
|
|
source = pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency, connection_id, remote_data) "
|
|
"VALUES ('alice', 'Drive', 'connector:file', 'weekly', CAST(:c AS uuid), "
|
|
"'{\"provider\": \"google_drive\"}') RETURNING id"
|
|
), {"c": cid}).scalar()
|
|
ConnectorPoliciesRepository(pg_conn).upsert("google_drive", enabled=False)
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), patch.object(worker, "ingest_connector") as ingest:
|
|
result = worker.sync_connector_source(MagicMock(), str(source))
|
|
assert result == {"status": "disabled"}
|
|
ingest.assert_not_called()
|
|
|
|
def test_disabled_connector_gives_remote_sync_no_credentials(self, pg_conn):
|
|
from docsgpt import worker
|
|
from docsgpt.storage.db.repositories.connector_policies import ConnectorPoliciesRepository
|
|
|
|
cid = _connection(pg_conn, provider="s3", auth_kind="api_key",
|
|
secrets={"credentials": {"aws_access_key_id": "AKIA", "aws_secret_access_key": "s"}})
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield):
|
|
assert worker._with_connection_credentials({"bucket": "b"}, cid)["aws_access_key_id"] == "AKIA"
|
|
ConnectorPoliciesRepository(pg_conn).upsert("s3", enabled=False)
|
|
assert worker._with_connection_credentials({"bucket": "b"}, cid) is None
|
|
|
|
def test_repository_url_gets_the_connections_token(self, pg_conn):
|
|
"""A GitHub source's loader input is a plain URL, not JSON."""
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="github", auth_kind="api_key",
|
|
secrets={"credentials": {"access_token": "github_pat_x"}})
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield):
|
|
data = worker._with_connection_credentials("https://github.com/acme/private", cid)
|
|
as_json = worker._with_connection_credentials('{"search_queries": ["x"]}', cid)
|
|
assert data == {"url": "https://github.com/acme/private", "access_token": "github_pat_x"}
|
|
assert json.loads(as_json)["access_token"] == "github_pat_x"
|
|
|
|
def test_oauth_connection_gives_its_current_access_token(self, pg_conn):
|
|
"""A GitHub App sign-in keeps an OAuth token, refreshed before use."""
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="github", auth_kind="oauth",
|
|
secrets={"token_info": {"access_token": "ghu_fresh", "refresh_token": "ghr_x"}})
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), _service_db(pg_conn), patch(
|
|
"docsgpt.connectors.service.get_valid_token_info", return_value={"access_token": "ghu_fresh"},
|
|
) as valid:
|
|
data = worker._with_connection_credentials({"repo_url": "acme/private"}, cid)
|
|
valid.assert_called_once_with(cid)
|
|
assert data == {"repo_url": "acme/private", "access_token": "ghu_fresh"}
|
|
|
|
def test_rejected_token_flags_the_connection(self, pg_conn):
|
|
"""A revoked token pauses the source for reconnect instead of failing every sync."""
|
|
from docsgpt import worker
|
|
from docsgpt.connectors.service import ConnectionUnavailable
|
|
from docsgpt.parser.remote.github_loader import GitHubTokenRejected
|
|
|
|
loader = MagicMock()
|
|
loader.load_data.side_effect = GitHubTokenRejected("revoked")
|
|
task = MagicMock()
|
|
task.request.retries = 1
|
|
with patch.object(worker.RemoteCreator, "create_loader", return_value=loader), patch.object(
|
|
worker, "_with_connection_credentials", return_value={"url": "acme/r", "access_token": "t"},
|
|
), patch.object(worker, "publish_user_event"), patch(
|
|
"docsgpt.connectors.service.mark_reconnect_needed",
|
|
) as flag:
|
|
with pytest.raises(ConnectionUnavailable):
|
|
worker.remote_worker(task, "acme/r", "repo", "alice", "github", connection_id="c-1")
|
|
flag.assert_called_once()
|
|
assert flag.call_args.args[0] == "c-1"
|
|
|
|
def test_sync_runs_as_the_connection_without_a_browser(self, pg_conn):
|
|
from docsgpt import worker
|
|
|
|
cid = _connection(pg_conn, provider="google_drive", auth_kind="oauth",
|
|
secrets={"token_info": {"access_token": "a"}})
|
|
source = pg_conn.execute(text(
|
|
"INSERT INTO sources (user_id, name, type, sync_frequency, connection_id, remote_data) "
|
|
"VALUES ('alice', 'Drive', 'connector:file', 'daily', CAST(:c AS uuid), "
|
|
"'{\"provider\": \"google_drive\", \"folder_ids\": [\"f\"], \"recursive\": false}') RETURNING id"
|
|
), {"c": cid}).scalar()
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
with patch.object(worker, "db_readonly", _yield), patch.object(
|
|
worker, "ingest_connector", return_value={"tokens": 1}
|
|
) as ingest:
|
|
result = worker.sync_connector_source(MagicMock(), str(source))
|
|
assert result["status"] == "success"
|
|
kwargs = ingest.call_args.kwargs
|
|
assert kwargs["connection_id"] == cid
|
|
assert kwargs["operation_mode"] == "sync"
|
|
assert kwargs["folder_ids"] == ["f"] and kwargs["recursive"] is False
|
|
assert "session_token" not in kwargs
|
|
|
|
|
|
@pytest.mark.parametrize("secret_key", ["token_info", "tokens", "client_info", "encrypted_credentials",
|
|
"client_secret", "refresh_token", "access_token"])
|
|
def test_redaction_covers_connection_secrets(secret_key):
|
|
from docsgpt.storage.db.redaction import REDACTED, redact_secrets
|
|
|
|
assert redact_secrets({secret_key: {"x": "y"}})[secret_key] == REDACTED
|
|
|
|
|
|
class TestMcpServerMismatch:
|
|
def test_connection_for_another_server_is_not_applied(self, pg_conn):
|
|
"""A key stored for one MCP server is never sent to a tool now pointing at another."""
|
|
from docsgpt.connectors import service
|
|
|
|
cid = _connection(pg_conn, provider="custom_mcp", server_url="https://old.example.com",
|
|
secrets={"credentials": {"bearer_token": "old-server-secret"}})
|
|
tool = {**_tool(cid, name="mcp_tool"), "config": {"server_url": "https://new.example.com/mcp",
|
|
"auth_type": "bearer"}}
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
with pytest.raises(service.ConnectionUnavailable):
|
|
_executor()._get_or_load_tool(tool, "t1", "search")
|
|
manager.return_value.load_tool.assert_not_called()
|
|
|
|
@pytest.mark.parametrize("tool_name, provider, server_url", [
|
|
# A service's key without a server of its own is not an MCP server's.
|
|
("mcp_tool", "telegram", None),
|
|
("mcp_tool", "ntfy", None),
|
|
# A custom server connection that names no server has nowhere to go.
|
|
("mcp_tool", "custom_mcp", None),
|
|
# A tool only runs on a connection of the connector that provides it.
|
|
("ntfy", "telegram", None),
|
|
("telegram", "ntfy", None),
|
|
("telegram", "custom_mcp", "https://new.example.com"),
|
|
])
|
|
def test_connection_of_another_connector_is_not_applied(self, pg_conn, tool_name, provider, server_url):
|
|
from docsgpt.connectors import service
|
|
|
|
cid = _connection(pg_conn, provider=provider, server_url=server_url,
|
|
secrets={"credentials": {"token": "bot-token", "bearer_token": "bot-token"}})
|
|
tool = {**_tool(cid, name=tool_name), "config": {"server_url": "https://new.example.com/mcp",
|
|
"auth_type": "bearer"}}
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
with pytest.raises(service.ConnectionUnavailable):
|
|
_executor()._get_or_load_tool(tool, "t1", "search")
|
|
manager.return_value.load_tool.assert_not_called()
|
|
|
|
def test_legacy_mcp_connection_is_still_applied_to_its_server(self, pg_conn):
|
|
"""Rows from before connector keys are named from their ``mcp:`` provider."""
|
|
cid = _connection(pg_conn, provider="mcp:https://m.example.com", auth_kind="mcp_oauth",
|
|
server_url=None, secrets={"tokens": {"access_token": "x"}})
|
|
pg_conn.execute(text("UPDATE connector_sessions SET connector_key = NULL WHERE id = CAST(:i AS uuid)"),
|
|
{"i": cid})
|
|
tool = {**_tool(cid, name="mcp_tool"), "config": {"server_url": "https://m.example.com/mcp",
|
|
"auth_type": "oauth"}}
|
|
with _service_db(pg_conn), patch("docsgpt.agents.tool_executor.ToolManager") as manager:
|
|
_executor()._get_or_load_tool(tool, "t1", "search")
|
|
assert manager.return_value.load_tool.call_args.kwargs["tool_config"]["connection_id"] == cid
|
|
|
|
def test_save_keeps_previous_connection_only_for_the_same_server(self, pg_conn):
|
|
from docsgpt.api.user.tools.mcp import _previous_connection
|
|
|
|
@contextmanager
|
|
def _yield():
|
|
yield pg_conn
|
|
|
|
cid = _connection(pg_conn, provider="custom_mcp", server_url="https://old.example.com")
|
|
existing = {"connection_id": cid}
|
|
same = {"server_url": "https://old.example.com/mcp", "auth_type": "bearer"}
|
|
with patch("docsgpt.api.user.tools.mcp.db_readonly", _yield):
|
|
assert _previous_connection(existing, same, "alice") == cid
|
|
assert _previous_connection(existing, {**same, "server_url": "https://new.example.com/mcp"}, "alice") is None
|
|
assert _previous_connection(None, same, "alice") is None
|
|
# Someone else's connection, or one that signs in another way, is not kept.
|
|
assert _previous_connection(existing, same, "bob") is None
|
|
assert _previous_connection(existing, {**same, "auth_type": "oauth"}, "alice") is None
|
|
|
|
|
|
class TestExternalApiCallers:
|
|
"""An agent called with its API key runs as the owner, and nobody can approve there."""
|
|
|
|
def _external(self, allowlist=None):
|
|
from docsgpt.agents.tool_executor import ToolExecutor
|
|
|
|
return ToolExecutor(user="alice", external_caller=True, api_write_allowlist=allowlist)
|
|
|
|
def test_owner_account_write_is_denied(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
pause = _pause(self._external(), _tool(cid))
|
|
assert pause["pause_type"] == "headless_denied"
|
|
assert pause["error_type"] == "tool_not_allowed"
|
|
assert "Access details" in pause["deny_reason"]
|
|
|
|
def test_even_always_allow_writes_are_denied(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
tool = _tool(cid)
|
|
tool["actions"][0]["require_approval"] = False
|
|
with _service_db(pg_conn):
|
|
assert _pause(self._external(), tool)["pause_type"] == "headless_denied"
|
|
|
|
def test_allowlisted_write_runs(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
pause = _pause(self._external(["tool-1:telegram_send_message"]), _tool(cid))
|
|
assert pause is None
|
|
|
|
def test_allowlist_does_not_cover_other_actions(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
pause = _pause(self._external(["tool-1:telegram_send_image"]), _tool(cid))
|
|
assert pause["pause_type"] == "headless_denied"
|
|
|
|
def test_missing_connection_is_denied_not_paused(self, pg_conn):
|
|
"""The widget cannot show a Connect card."""
|
|
cid = _connection(pg_conn, status="reconnect_needed")
|
|
with _service_db(pg_conn):
|
|
pause = _pause(self._external(), _tool(cid))
|
|
assert pause["pause_type"] == "headless_denied"
|
|
assert pause["error_type"] == "connection_required"
|
|
|
|
def test_the_owner_in_the_app_is_not_external(self, pg_conn):
|
|
cid = _connection(pg_conn)
|
|
with _service_db(pg_conn):
|
|
assert _pause(_executor(), _tool(cid)) is None
|
|
|
|
|
|
class TestExternalCallerDetection:
|
|
def test_api_key_request_from_someone_else_is_external(self):
|
|
from docsgpt.api.answer.services.stream_processor import is_external_api_caller
|
|
|
|
assert is_external_api_caller({"api_key": "k"}, {"sub": "visitor"}, "alice") is True
|
|
assert is_external_api_caller({"api_key": "k"}, None, "alice") is True
|
|
|
|
def test_owner_previewing_their_agent_is_not_external(self):
|
|
from docsgpt.api.answer.services.stream_processor import is_external_api_caller
|
|
|
|
assert is_external_api_caller({"api_key": "k"}, {"sub": "alice"}, "alice") is False
|
|
assert is_external_api_caller({}, {"sub": "visitor"}, "alice") is False
|
|
|
|
|
|
class TestApiWriteAllowlistConfig:
|
|
def test_accepts_tool_action_pairs(self):
|
|
from docsgpt.guardrails.config import AgentConfig
|
|
|
|
config = AgentConfig.model_validate({"api_write_allowlist": ["tool-1:telegram_send_message"]})
|
|
assert config.api_write_allowlist == ["tool-1:telegram_send_message"]
|
|
|
|
def test_rejects_malformed_entries(self):
|
|
from docsgpt.guardrails.config import AgentConfig
|
|
|
|
with pytest.raises(Exception):
|
|
AgentConfig.model_validate({"api_write_allowlist": ["no-action-part"]})
|
|
|
|
def test_old_configs_still_parse(self):
|
|
from docsgpt.guardrails.config import AgentConfig
|
|
|
|
assert AgentConfig.parse({"guardrails": {}}).api_write_allowlist == []
|
|
|
|
|
|
class TestAllowlistOwnership:
|
|
def test_team_editor_cannot_change_the_allowlist(self):
|
|
from docsgpt.api.user.agents.routes import keep_owner_only_config
|
|
|
|
existing = {"config": {"api_write_allowlist": ["t:a"]}}
|
|
sent = {"guardrails": {"controls": []}, "api_write_allowlist": ["t:a", "t:b"]}
|
|
assert keep_owner_only_config(sent, existing, True)["api_write_allowlist"] == ["t:a"]
|
|
assert keep_owner_only_config(sent, existing, False)["api_write_allowlist"] == ["t:a", "t:b"]
|
|
assert keep_owner_only_config({}, {"config": None}, True) == {"api_write_allowlist": []}
|