mirror of
https://github.com/tiennm99/DocsGPT.git
synced 2026-10-11 12:11:45 +00:00
Let people name accounts and use the names to tell accounts apart
Connections get an account_name (migration 0039) that the owner sets with PATCH /api/connections/<id>; the account label stays the account's identity, so signing in again still finds it. When someone has more than one account of a service, its tools are listed as "Telegram · <account>" and the model sees each account's actions under the account's name (telegram_send_message_alerts_bot) with the account in the description, instead of _1 and _2. Actions shared by different services are named after the service. Tools a user renamed keep their name.
This commit is contained in:
1 parent
607c2296e3
commit
bdcb88866a
11 files changed
+371
-7
No files matched your search
@@ -149,6 +149,16 @@ def _requires_approval(tool: Dict, action: Dict) -> bool:
|
||||
return bool((tool.get("config") or {}).get("require_approval"))
|
||||
|
||||
|
||||
def _account_slug(account: Optional[str], limit: int = 24) -> str:
|
||||
"""An account name as a function-name suffix: ``Ops: on-call!`` → ``ops_on_call``.
|
||||
|
||||
Empty when nothing ASCII is left (a name in another script); the caller
|
||||
then numbers the duplicates instead.
|
||||
"""
|
||||
slug = re.sub(r"[^a-z0-9]+", "_", str(account or "").lower()).strip("_")
|
||||
return slug[:limit].rstrip("_")
|
||||
|
||||
|
||||
def _sanitize_tool_prefix(tool_name: Optional[str]) -> str:
|
||||
"""Reduce a tool name to characters allowed in function-call names."""
|
||||
return re.sub(r"[^a-zA-Z0-9_-]+", "_", str(tool_name or "")).strip("_")
|
||||
@@ -723,15 +733,38 @@ class ToolExecutor:
|
||||
self._tool_to_name = {}
|
||||
all_llm_names: set = set()
|
||||
|
||||
# Connection tools that share an action name are told apart by what
|
||||
# they connect to: the service ("search" on Notion and on Linear), or
|
||||
# the account when one service is connected twice (two Telegram bots).
|
||||
connected: Dict[int, Tuple[str, str]] = {}
|
||||
for index, (tool_id, _tool_name, action_name, _action, is_client) in enumerate(entries):
|
||||
if name_counts[action_name] > 1 and not is_client:
|
||||
names = self._connection_names(tools_dict[tool_id])
|
||||
if names:
|
||||
connected[index] = names
|
||||
per_service = Counter((entries[i][2], service) for i, (service, _account) in connected.items())
|
||||
|
||||
result = []
|
||||
for tool_id, tool_name, action_name, action, is_client in entries:
|
||||
for index, (tool_id, tool_name, action_name, action, is_client) in enumerate(entries):
|
||||
service, account = connected.get(index, (None, None))
|
||||
if service is None:
|
||||
slug = ""
|
||||
elif per_service[(action_name, service)] > 1:
|
||||
slug = _account_slug(account)
|
||||
else:
|
||||
slug = _account_slug(service)
|
||||
if name_counts[action_name] == 1 and len(action_name) <= _MAX_LLM_NAME_LEN:
|
||||
llm_name = action_name
|
||||
else:
|
||||
# An over-long unique name skips the prefix — it needs
|
||||
# truncation, not disambiguation.
|
||||
prefix = _sanitize_tool_prefix(tool_name) if name_counts[action_name] > 1 else ""
|
||||
base = f"{prefix}_{action_name}" if prefix and not action_name.startswith(f"{prefix}_") else action_name
|
||||
if slug:
|
||||
base = f"{action_name}_{slug}"
|
||||
elif prefix and not action_name.startswith(f"{prefix}_"):
|
||||
base = f"{prefix}_{action_name}"
|
||||
else:
|
||||
base = action_name
|
||||
base = base[:_MAX_LLM_NAME_LEN]
|
||||
# A duplicated bare name stays ambiguous, and a candidate
|
||||
# must not steal a unique action's name or one already taken.
|
||||
@@ -754,18 +787,32 @@ class ToolExecutor:
|
||||
action, hidden=set(self._connection_parameters(tools_dict[tool_id])),
|
||||
)
|
||||
|
||||
description = action.get("description", "")
|
||||
if service:
|
||||
description = f"{description} ({service} account: {account})".strip()
|
||||
result.append(
|
||||
{
|
||||
"type": "function",
|
||||
"function": {
|
||||
"name": llm_name,
|
||||
"description": action.get("description", ""),
|
||||
"description": description,
|
||||
"parameters": params,
|
||||
},
|
||||
}
|
||||
)
|
||||
return result
|
||||
|
||||
def _connection_names(self, tool_data: Dict) -> Optional[Tuple[str, str]]:
|
||||
"""``(service, account)`` a connection tool runs with, e.g. ``("Telegram", "Alerts bot")``."""
|
||||
if not tool_data.get("connection_id") or tool_data.get("client_side"):
|
||||
return None
|
||||
resolved = self._resolve_connection(tool_data)
|
||||
if resolved is None or resolved.row is None:
|
||||
return None
|
||||
from docsgpt.connectors.service import account_name
|
||||
|
||||
return resolved.connector_name or tool_data.get("name") or "", account_name(resolved.row)
|
||||
|
||||
def _build_tool_parameters(self, action: Dict, hidden: Optional[set] = None) -> Dict:
|
||||
"""The JSON schema the model sees for ``action``.
|
||||
|
||||
|
||||
@@ -0,0 +1,32 @@
|
||||
"""0039 connection account name — what the user calls an account.
|
||||
|
||||
``account_label`` identifies an account: the email an OAuth sign-in returns,
|
||||
or a hint of a pasted key. Signing in again finds the connection by it, so it
|
||||
cannot be renamed. ``account_name`` is the name the user gives the account
|
||||
("Alerts bot", "Work"), shown instead of the label and used to tell two
|
||||
accounts of one service apart, for people and for the model. NULL means the
|
||||
user never named it.
|
||||
|
||||
Idempotent both ways.
|
||||
|
||||
Revision ID: 0039_connection_account_name
|
||||
Revises: 0038_connections
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
|
||||
|
||||
revision: str = "0039_connection_account_name"
|
||||
down_revision: Union[str, None] = "0038_connections"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.execute("ALTER TABLE connector_sessions ADD COLUMN IF NOT EXISTS account_name TEXT;")
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.execute("ALTER TABLE connector_sessions DROP COLUMN IF EXISTS account_name;")
|
||||
@@ -161,6 +161,21 @@ class ConnectionDetail(Resource):
|
||||
return make_response(jsonify({"success": False, "error": "Failed to load connection"}), 500)
|
||||
return make_response(jsonify({"success": True, "connection": detail}), 200)
|
||||
|
||||
@api.doc(description="Name an account: {name}. An empty name clears it. Owner only.")
|
||||
def patch(self, connection_id: str):
|
||||
user_id = _user_id()
|
||||
if not user_id:
|
||||
return _unauthorized()
|
||||
name = _json_body().get("name")
|
||||
if not isinstance(name, str) or len(name.strip()) > service.ACCOUNT_NAME_MAX:
|
||||
return _error(f"name must be text of at most {service.ACCOUNT_NAME_MAX} characters", 400)
|
||||
with db_session() as conn:
|
||||
row = _owned(conn, connection_id, user_id)
|
||||
if row is None:
|
||||
return _not_found()
|
||||
connection = service.rename_connection(conn, row, name)
|
||||
return make_response(jsonify({"success": True, "connection": connection}), 200)
|
||||
|
||||
@api.doc(
|
||||
description=(
|
||||
"Remove a connection: {sources: keep | delete, tools: delete | keep}. "
|
||||
|
||||
@@ -23,6 +23,7 @@ from docsgpt.api.pat.rules import filter_listing
|
||||
from docsgpt.api.user.artifacts.authz import Principal, authorize_artifact
|
||||
from docsgpt.api.user.team_sharing import effective_write_owner, visible_with_access
|
||||
from docsgpt.connectors.catalog import definition_for_tool
|
||||
from docsgpt.connectors.service import account_tool_names
|
||||
from docsgpt.core.settings import settings
|
||||
from docsgpt.core.url_validation import SSRFError, validate_url
|
||||
from docsgpt.security.encryption import CredentialDecryptionError, decrypt_credentials, encrypt_credentials
|
||||
@@ -290,6 +291,8 @@ class GetTools(Resource):
|
||||
team_shared = visible_with_access(conn, user, "tool")
|
||||
shared_ids = [tid for tid in team_shared if tid not in owned_ids]
|
||||
shared_rows = tools_repo.list_by_ids(shared_ids)
|
||||
# "Telegram · Alerts bot" when the owner has several bots.
|
||||
account_names = account_tool_names(conn, [*rows, *shared_rows])
|
||||
user_tools = []
|
||||
|
||||
def _shape_tool(row, *, ownership="user", force_strip_secret=False):
|
||||
@@ -313,6 +316,8 @@ class GetTools(Resource):
|
||||
# ask for it again.
|
||||
tool_copy.setdefault("config", {})["has_encrypted_credentials"] = True
|
||||
tool_copy["ownership"] = ownership
|
||||
if str(row["id"]) in account_names:
|
||||
tool_copy["customName"] = tool_copy["displayName"] = account_names[str(row["id"])]
|
||||
return tool_copy
|
||||
|
||||
for row in rows:
|
||||
|
||||
@@ -96,6 +96,78 @@ def account_label(row: dict) -> str:
|
||||
)
|
||||
|
||||
|
||||
#: Longest name a user can give an account.
|
||||
ACCOUNT_NAME_MAX = 80
|
||||
|
||||
|
||||
def account_name(row: dict) -> str:
|
||||
"""What the user calls the account, falling back to its label."""
|
||||
return (row.get("account_name") or "").strip() or account_label(row)
|
||||
|
||||
|
||||
def rename_connection(conn, row: dict, name: str) -> dict:
|
||||
"""Give an account a name (an empty one clears it).
|
||||
|
||||
The name is only a name: the account is still found by its label when
|
||||
its owner signs in again.
|
||||
|
||||
Args:
|
||||
conn: Open connection inside a transaction.
|
||||
row: The connection, already authorised for its owner.
|
||||
name: The new name; surrounding spaces are dropped.
|
||||
|
||||
Returns:
|
||||
The connection's public shape after the change.
|
||||
"""
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
repo.update(str(row["id"]), {"account_name": name.strip() or None})
|
||||
return serialize_connection(repo.get(str(row["id"])))
|
||||
|
||||
|
||||
def with_account(name: str, account: str) -> str:
|
||||
"""``Telegram · Alerts bot``: a tool's name with the account it uses."""
|
||||
return f"{name} · {account}" if account and account not in name else name
|
||||
|
||||
|
||||
def account_tool_names(conn, tools: Iterable[dict]) -> dict[str, str]:
|
||||
"""Names for tools whose owner has several accounts of the tool's service.
|
||||
|
||||
With one Telegram bot a tool is just "Telegram"; with two, each tool
|
||||
still named after the service becomes "Telegram · <account>", so people
|
||||
(and the model) can tell them apart. A name the user chose is kept.
|
||||
|
||||
Args:
|
||||
conn: Open database connection.
|
||||
tools: ``user_tools`` rows.
|
||||
|
||||
Returns:
|
||||
Tool id to its name, only for tools whose name changes.
|
||||
"""
|
||||
repo = ConnectorSessionsRepository(conn)
|
||||
connections: dict[str, dict] = {}
|
||||
for tool in tools:
|
||||
connection_id = str(tool.get("connection_id") or "")
|
||||
if connection_id and connection_id not in connections:
|
||||
row = repo.get(connection_id)
|
||||
if row is not None:
|
||||
connections[connection_id] = row
|
||||
counts: dict[tuple, int] = {}
|
||||
for owner in {row["user_id"] for row in connections.values()}:
|
||||
for row in repo.list_for_user(owner):
|
||||
if normalize_status(row) != STATUS_PENDING:
|
||||
key = (owner, catalog.connector_key_for_row(row))
|
||||
counts[key] = counts.get(key, 0) + 1
|
||||
names = {}
|
||||
for tool in tools:
|
||||
row = connections.get(str(tool.get("connection_id") or ""))
|
||||
if row is None or counts.get((row["user_id"], catalog.connector_key_for_row(row)), 0) < 2:
|
||||
continue
|
||||
name = tool.get("custom_name") or tool.get("display_name") or ""
|
||||
if name == serialize_connection(row)["name"]:
|
||||
names[str(tool["id"])] = with_account(name, account_name(row))
|
||||
return names
|
||||
|
||||
|
||||
def _iso(value: Any) -> Optional[str]:
|
||||
return value.isoformat() if hasattr(value, "isoformat") else value
|
||||
|
||||
@@ -114,6 +186,7 @@ def serialize_connection(row: dict, counts: Optional[dict] = None) -> dict:
|
||||
"display_name": row.get("display_name"),
|
||||
"icon": definition.icon if definition else "tool_mcp_tool",
|
||||
"account_label": account_label(row),
|
||||
"account_name": (row.get("account_name") or "").strip() or None,
|
||||
"auth_kind": row.get("auth_kind") or (definition.auth_kind if definition else None),
|
||||
"status": normalize_status(row),
|
||||
"server_url": row.get("server_url"),
|
||||
@@ -181,13 +254,17 @@ def serialize_parameters(action: dict, account_parameters: Optional[dict] = None
|
||||
return parameters
|
||||
|
||||
|
||||
def serialize_tool(row: dict, account_parameters: Optional[dict] = None) -> dict:
|
||||
def serialize_tool(
|
||||
row: dict, account_parameters: Optional[dict] = None, display_name: Optional[str] = None,
|
||||
) -> dict:
|
||||
"""A tool linked to a connection, with its actions, permissions and parameters.
|
||||
|
||||
Args:
|
||||
row: The ``user_tools`` row.
|
||||
account_parameters: Parameters its connection sets, from
|
||||
:func:`connection_parameters`.
|
||||
display_name: The name to show instead of the stored one, from
|
||||
:func:`account_tool_names`.
|
||||
"""
|
||||
from docsgpt.connectors.permissions import action_access, action_permission
|
||||
|
||||
@@ -208,7 +285,7 @@ def serialize_tool(row: dict, account_parameters: Optional[dict] = None) -> dict
|
||||
return {
|
||||
"id": str(row["id"]),
|
||||
"name": row.get("name"),
|
||||
"display_name": row.get("custom_name") or row.get("display_name") or row.get("name"),
|
||||
"display_name": display_name or row.get("custom_name") or row.get("display_name") or row.get("name"),
|
||||
"status": bool(row.get("status")),
|
||||
"credential_mode": row.get("credential_mode") or "owner",
|
||||
"actions": actions,
|
||||
@@ -221,7 +298,9 @@ def connection_detail(conn, row: dict) -> dict:
|
||||
status = normalize_status(row)
|
||||
sources = [serialize_source(s, status) for s in repo.list_sources(str(row["id"]))]
|
||||
account_parameters = connection_parameters(row) if status == STATUS_CONNECTED else {}
|
||||
tools = [serialize_tool(t, account_parameters) for t in repo.list_tools(str(row["id"]))]
|
||||
tool_rows = repo.list_tools(str(row["id"]))
|
||||
names = account_tool_names(conn, tool_rows)
|
||||
tools = [serialize_tool(t, account_parameters, names.get(str(t["id"]))) for t in tool_rows]
|
||||
detail = serialize_connection(row, {"sources": len(sources), "tools": len(tools)})
|
||||
detail["sources"] = sources
|
||||
detail["tools"] = tools
|
||||
|
||||
@@ -608,6 +608,9 @@ connector_sessions_table = Table(
|
||||
Column("scopes", JSONB, nullable=False, server_default="[]"),
|
||||
Column("last_error", Text),
|
||||
Column("last_used_at", DateTime(timezone=True)),
|
||||
# Added in ``0039_connection_account_name``: what the user calls the
|
||||
# account; ``account_label`` stays its identity.
|
||||
Column("account_name", Text),
|
||||
)
|
||||
|
||||
|
||||
|
||||
@@ -36,7 +36,7 @@ from docsgpt.storage.db.serialization import PGNativeJSONEncoder
|
||||
_UPDATABLE_SCALARS = {
|
||||
"server_url", "session_token", "user_email", "status", "expires_at",
|
||||
"connector_key", "display_name", "account_label", "auth_kind",
|
||||
"encrypted_credentials", "has_refresh_token", "last_error", "last_used_at",
|
||||
"encrypted_credentials", "has_refresh_token", "last_error", "last_used_at", "account_name",
|
||||
}
|
||||
_UPDATABLE_JSONB = {"session_data", "token_info", "scopes"}
|
||||
|
||||
|
||||
@@ -417,6 +417,41 @@ class TestParameters:
|
||||
assert chat_id["filled_by_llm"] is True
|
||||
|
||||
|
||||
class TestRename:
|
||||
@staticmethod
|
||||
def _patch(app, cid, body, user="alice"):
|
||||
from docsgpt.api.connector.connections import ConnectionDetail
|
||||
|
||||
return _call(app, ConnectionDetail, "patch", f"/api/connections/{cid}", user=user, body=body, args=[cid])
|
||||
|
||||
def test_owner_names_and_unnames_an_account(self, app, pg_conn):
|
||||
cid = _connection(pg_conn, secrets={"credentials": {"token": "t"}})
|
||||
with _db(pg_conn):
|
||||
named = self._patch(app, cid, {"name": " Alerts bot "})
|
||||
assert named.status_code == 200
|
||||
connection = named.get_json()["connection"]
|
||||
assert connection["account_name"] == "Alerts bot"
|
||||
# The label stays the account's identity.
|
||||
assert connection["account_label"] == "…abcd"
|
||||
cleared = self._patch(app, cid, {"name": ""})
|
||||
assert cleared.get_json()["connection"]["account_name"] is None
|
||||
|
||||
@pytest.mark.parametrize("body", [{"name": "x" * 81}, {"name": 5}, {}])
|
||||
def test_rejects_bad_names(self, app, pg_conn, body):
|
||||
cid = _connection(pg_conn, secrets={})
|
||||
with _db(pg_conn):
|
||||
assert self._patch(app, cid, body).status_code == 400
|
||||
|
||||
def test_another_user_cannot_rename(self, app, pg_conn):
|
||||
cid = _connection(pg_conn, secrets={})
|
||||
with _db(pg_conn):
|
||||
assert self._patch(app, cid, {"name": "Mine now"}, user="mallory").status_code == 404
|
||||
name = pg_conn.execute(
|
||||
text("SELECT account_name FROM connector_sessions WHERE id = CAST(:i AS uuid)"), {"i": cid}
|
||||
).scalar()
|
||||
assert name is None
|
||||
|
||||
|
||||
class TestDelete:
|
||||
def test_delete_with_source_removal(self, app, pg_conn):
|
||||
from docsgpt.api.connector.connections import ConnectionDetail
|
||||
|
||||
@@ -276,6 +276,55 @@ class TestTelegramDefaultChat:
|
||||
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"].startswith("Search (Wiki account:")
|
||||
|
||||
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
|
||||
|
||||
@@ -245,3 +245,52 @@ class TestAvailableTools:
|
||||
assert tools["telegram"]["displayName"] == "Telegram"
|
||||
assert tools["ntfy"]["displayName"] == "ntfy"
|
||||
assert tools["postgres"]["displayName"] == "PostgreSQL"
|
||||
|
||||
|
||||
def _telegram_connection(conn, user, label, name=None):
|
||||
row = conn.execute(
|
||||
text(
|
||||
"INSERT INTO connector_sessions (user_id, provider, connector_key, auth_kind, status, account_label, "
|
||||
"account_name, encrypted_credentials) VALUES (:u, 'telegram', 'telegram', 'api_key', 'connected', "
|
||||
":l, :n, :e) RETURNING *"
|
||||
),
|
||||
{"u": user, "l": label, "n": name, "e": encrypt_json({"credentials": {"token": label}}, user)},
|
||||
).one()
|
||||
return dict(row._mapping)
|
||||
|
||||
|
||||
class TestAccountNamesInToolNames:
|
||||
def _listed(self, app, pg_conn, user="alice"):
|
||||
from docsgpt.api.user.tools.routes import GetTools
|
||||
|
||||
with _db(pg_conn), app.test_request_context("/api/get_tools"):
|
||||
from flask import request
|
||||
|
||||
request.decoded_token = {"sub": user}
|
||||
tools = GetTools().get().get_json()["tools"]
|
||||
return sorted(t["customName"] for t in tools if t.get("name") == "telegram")
|
||||
|
||||
def test_two_accounts_are_named_after_their_accounts(self, app, pg_conn):
|
||||
for label, name in (("…aaaa", "Alerts bot"), ("…bbbb", None)):
|
||||
service.ensure_connection_tools(pg_conn, "alice", _telegram_connection(pg_conn, "alice", label, name))
|
||||
assert self._listed(app, pg_conn) == ["Telegram · Alerts bot", "Telegram · …bbbb"]
|
||||
|
||||
def test_one_account_keeps_the_plain_name(self, app, pg_conn):
|
||||
connection = _telegram_connection(pg_conn, "alice", "…aaaa", "Alerts bot")
|
||||
service.ensure_connection_tools(pg_conn, "alice", connection)
|
||||
assert self._listed(app, pg_conn) == ["Telegram"]
|
||||
|
||||
def test_a_name_the_user_chose_is_kept(self, app, pg_conn):
|
||||
for label in ("…aaaa", "…bbbb"):
|
||||
service.ensure_connection_tools(pg_conn, "alice", _telegram_connection(pg_conn, "alice", label))
|
||||
pg_conn.execute(text("UPDATE user_tools SET custom_name = 'Ops' WHERE name = 'telegram' "
|
||||
"AND connection_id = (SELECT id FROM connector_sessions WHERE account_label = '…aaaa')"))
|
||||
assert self._listed(app, pg_conn) == ["Ops", "Telegram · …bbbb"]
|
||||
|
||||
def test_the_drawer_names_the_tool_after_its_account_too(self, pg_conn):
|
||||
named = _telegram_connection(pg_conn, "alice", "…aaaa", "Alerts bot")
|
||||
for connection in (named, _telegram_connection(pg_conn, "alice", "…bbbb")):
|
||||
service.ensure_connection_tools(pg_conn, "alice", connection)
|
||||
detail = service.connection_detail(pg_conn, named)
|
||||
assert detail["account_name"] == "Alerts bot"
|
||||
assert detail["tools"][0]["display_name"] == "Telegram · Alerts bot"
|
||||
@@ -0,0 +1,50 @@
|
||||
"""Migration round-trip test for 0039_connection_account_name."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
from pathlib import Path
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import text
|
||||
|
||||
|
||||
pytestmark = pytest.mark.integration
|
||||
|
||||
_0039 = "0039_connection_account_name"
|
||||
_0038 = "0038_connections"
|
||||
|
||||
|
||||
def _run_alembic(url: str, *args: str) -> None:
|
||||
ini = Path(__file__).resolve().parents[3] / "docsgpt" / "alembic.ini"
|
||||
subprocess.check_call(
|
||||
[sys.executable, "-m", "alembic", "-c", str(ini), *args],
|
||||
timeout=120,
|
||||
env={**os.environ, "POSTGRES_URI": url},
|
||||
)
|
||||
|
||||
|
||||
def _has_column(conn) -> bool:
|
||||
return conn.execute(
|
||||
text(
|
||||
"SELECT 1 FROM information_schema.columns WHERE table_schema = 'public' "
|
||||
"AND table_name = 'connector_sessions' AND column_name = 'account_name'"
|
||||
)
|
||||
).fetchone() is not None
|
||||
|
||||
|
||||
class TestMigration0039RoundTrip:
|
||||
def test_head_has_account_name(self, pg_engine):
|
||||
with pg_engine.connect() as conn:
|
||||
assert _has_column(conn)
|
||||
|
||||
def test_downgrade_drops_then_upgrade_restores(self, pg_engine):
|
||||
url = pg_engine.url.render_as_string(hide_password=False)
|
||||
_run_alembic(url, "downgrade", _0038)
|
||||
with pg_engine.connect() as conn:
|
||||
assert not _has_column(conn)
|
||||
_run_alembic(url, "upgrade", "head")
|
||||
with pg_engine.connect() as conn:
|
||||
assert _has_column(conn)
|
||||
Reference in new issue
Block a user