test_workspace_provision.py 10 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191192193194195196197198199200201202203204205206207208209210211212213214215216217218219220221222223224225226227228229230231232233234235236237238239240241242
  1. """
  2. tests/test_workspace_provision.py — 创建智能体实例时自动开辟工作目录。
  3. 布局参照 Workspace/multi-reviewer67:
  4. <base>/<名>/ work_dir(最终交付物 + Bash CWD)
  5. <base>/<名>/data source_dir
  6. <base>/<名>/workspace run_dir(bare 模式,run_* 直接建在其下)
  7. """
  8. from __future__ import annotations
  9. import json
  10. import os
  11. import pytest
  12. os.environ.setdefault("AGENTPAAS_DATABASE_URL", "sqlite:///:memory:")
  13. os.environ.setdefault("AGENTPAAS_TESTING", "1")
  14. # ── 单元:provision 函数 ─────────────────────────────────────────────────────
  15. def test_provision_layout(tmp_path):
  16. from agentpaas.engine.workspace_provision import provision_agent_workspace
  17. d = provision_agent_workspace("审稿助手", base=str(tmp_path))
  18. root = os.path.join(str(tmp_path), "审稿助手")
  19. assert d["work_dir"] == os.path.abspath(root)
  20. assert d["source_dir"] == os.path.join(os.path.abspath(root), "data")
  21. assert d["run_dir"] == os.path.join(os.path.abspath(root), "workspace")
  22. assert os.path.isdir(d["source_dir"]) and os.path.isdir(d["run_dir"])
  23. def test_provision_duplicate_names_get_suffix(tmp_path):
  24. from agentpaas.engine.workspace_provision import provision_agent_workspace
  25. d1 = provision_agent_workspace("助手", base=str(tmp_path))
  26. d2 = provision_agent_workspace("助手", base=str(tmp_path))
  27. d3 = provision_agent_workspace("助手", base=str(tmp_path))
  28. assert d1["work_dir"].endswith("助手")
  29. assert d2["work_dir"].endswith("助手-2")
  30. assert d3["work_dir"].endswith("助手-3")
  31. def test_provision_slug_sanitises_path_chars(tmp_path):
  32. from agentpaas.engine.workspace_provision import provision_agent_workspace
  33. d = provision_agent_workspace("../evil/name v2", base=str(tmp_path))
  34. # 不逃逸 base;路径分隔符被清洗
  35. assert os.path.commonpath([d["work_dir"], str(tmp_path)]) == str(
  36. os.path.abspath(str(tmp_path)))
  37. assert "/evil/" not in d["work_dir"].replace(str(tmp_path), "")
  38. def test_provision_failure_returns_empty(monkeypatch):
  39. from agentpaas.engine import workspace_provision as wp
  40. monkeypatch.setattr(wp.os, "makedirs",
  41. lambda *a, **k: (_ for _ in ()).throw(OSError("ro")))
  42. d = wp.provision_agent_workspace("x", base="/nonexistent-base")
  43. assert d == {"work_dir": "", "source_dir": "", "run_dir": ""}
  44. # ── 集成:两个创建端点 ───────────────────────────────────────────────────────
  45. @pytest.fixture()
  46. def api_client(tmp_path, monkeypatch):
  47. """TestClient + 内存 DB + 租户/key + workspace_base 指到 tmp。"""
  48. import secrets
  49. from fastapi.testclient import TestClient
  50. from agentpaas.api.app import app
  51. from agentpaas.api.middleware.auth import hash_key
  52. from agentpaas.config import settings
  53. from agentpaas.db.models import Database, gen_id, now_utc
  54. import agentpaas.db.session as _session_mod
  55. ws_base = tmp_path / "Workspace"
  56. monkeypatch.setattr(settings, "workspace_base", str(ws_base))
  57. prev = _session_mod._db
  58. _session_mod._db = Database("sqlite:///:memory:")
  59. db = _session_mod._db
  60. tid, uid = gen_id("tn_"), gen_id("usr_")
  61. raw_key = f"ap_{secrets.token_hex(16)}"
  62. now = now_utc()
  63. db.execute("INSERT INTO tenants (id, name, plan, status, created_at) "
  64. "VALUES (?, 'test', 'free', 'active', ?)", (tid, now))
  65. db.execute("INSERT INTO users (id, tenant_id, email, role, created_at) "
  66. "VALUES (?, ?, '', 'admin', ?)", (uid, tid, now))
  67. db.execute(
  68. "INSERT INTO api_keys (id, tenant_id, user_id, key_hash, key_prefix, name, "
  69. "scopes, rate_limit, status, created_at) "
  70. "VALUES (?, ?, ?, ?, ?, 'test', ?, 600, 'active', ?)",
  71. (gen_id("key_"), tid, uid, hash_key(raw_key), raw_key[:8],
  72. json.dumps(["agents:*"]), now))
  73. db.commit()
  74. with TestClient(app) as client:
  75. yield client, raw_key, str(ws_base), db
  76. _session_mod._db = prev
  77. def _auth(k):
  78. return {"Authorization": f"Bearer {k}"}
  79. _MIN_CONFIG = {"name": "目录测试", "type": "simple",
  80. "model": {"name": "ollama/qwen2.5:7b"}, "systemPrompt": "t"}
  81. def test_create_agent_auto_provisions_dirs(api_client):
  82. client, key, ws_base, db = api_client
  83. r = client.post("/api/v1/agents", headers=_auth(key),
  84. json={"name": "目录测试", "config": _MIN_CONFIG})
  85. assert r.status_code in (200, 201), r.text[:300]
  86. body = r.json()
  87. assert body["work_dir"].endswith("目录测试")
  88. assert body["run_dir"] == os.path.join(body["work_dir"], "workspace")
  89. assert os.path.isdir(body["source_dir"])
  90. # DB 行一致
  91. row = db.fetchone("SELECT work_dir, source_dir, run_dir FROM agents WHERE id=?",
  92. (body["agent_id"],))
  93. assert row["work_dir"] == body["work_dir"]
  94. assert row["run_dir"] == body["run_dir"]
  95. def test_create_agent_respects_explicit_dirs(api_client, tmp_path):
  96. client, key, ws_base, db = api_client
  97. my_dir = str(tmp_path / "my-own-dir")
  98. os.makedirs(my_dir)
  99. r = client.post("/api/v1/agents", headers=_auth(key),
  100. json={"name": "显式目录", "config": _MIN_CONFIG,
  101. "work_dir": my_dir})
  102. assert r.status_code in (200, 201), r.text[:300]
  103. body = r.json()
  104. assert body["work_dir"] == my_dir
  105. assert body["run_dir"] == "" # 显式给了任一目录 → 其余不自动补
  106. # 没有为它建 Workspace/<名> 目录
  107. assert not os.path.exists(os.path.join(ws_base, "显式目录"))
  108. # ── 模型显示与切换(模型页功能)─────────────────────────────────────────────
  109. def test_models_in_use_and_switch(api_client):
  110. client, key, ws_base, db = api_client
  111. # 建两个不同模型的智能体
  112. for name, model in [("甲", {"provider": "ollama", "name": "qwen2.5:7b"}),
  113. ("乙", {"provider": "dashscope", "name": "qwen-max"})]:
  114. cfg = {"name": name, "type": "simple", "model": model, "systemPrompt": "t"}
  115. r = client.post("/api/v1/agents", headers=_auth(key),
  116. json={"name": name, "config": cfg})
  117. assert r.status_code in (200, 201), r.text[:200]
  118. # 一览
  119. r = client.get("/api/v1/providers/models-in-use", headers=_auth(key))
  120. assert r.status_code == 200, r.text[:200]
  121. items = {i["agent_name"]: i for i in r.json()["items"]}
  122. assert items["甲"]["provider"] == "ollama" and items["甲"]["model"] == "qwen2.5:7b"
  123. assert items["乙"]["provider"] == "dashscope" and items["乙"]["model"] == "qwen-max"
  124. # 切换甲 → dashscope/qwen-plus
  125. aid = items["甲"]["agent_id"]
  126. r = client.put(f"/api/v1/agents/{aid}/model", headers=_auth(key),
  127. json={"provider": "dashscope", "name": "qwen-plus"})
  128. assert r.status_code == 200, r.text[:300]
  129. body = r.json()
  130. assert body["version"] == 2 and body["previous"] == "ollama/qwen2.5:7b"
  131. # 配置生效且其余字段保留;版本历史可回滚
  132. g = client.get(f"/api/v1/agents/{aid}", headers=_auth(key)).json()
  133. assert g["config"]["model"]["provider"] == "dashscope"
  134. assert g["config"]["model"]["name"] == "qwen-plus"
  135. assert g["config"]["systemPrompt"] == "t"
  136. assert g["current_version"] == 2
  137. rb = client.post(f"/api/v1/agents/{aid}/rollback", headers=_auth(key),
  138. json={"target_version": 1})
  139. assert rb.status_code == 200
  140. g2 = client.get(f"/api/v1/agents/{aid}", headers=_auth(key)).json()
  141. assert g2["config"]["model"]["name"] == "qwen2.5:7b"
  142. def test_switch_model_unknown_agent_404(api_client):
  143. client, key, *_ = api_client
  144. r = client.put("/api/v1/agents/ag_nonexistent/model", headers=_auth(key),
  145. json={"provider": "ollama", "name": "x"})
  146. assert r.status_code == 404
  147. # ── simple agent token 透传(P3)────────────────────────────────────────────
  148. def test_simple_agent_usage_source_found():
  149. """_find_usage_source 沿包装链找到 get_usage:裸 ConversationLam、
  150. Guard 包装、Memory(Guard(...)) 双层包装都要命中。"""
  151. from lambdagent.conversation import ConversationLam
  152. from lambdagent.extensions import Guard, Memory
  153. class FakeProvider:
  154. default_model = "qwen-max"
  155. model_name = "qwen-max"
  156. context_window = 100000
  157. def chat(self, messages, **kw):
  158. return "ok"
  159. def get_usage(self):
  160. return {"input_tokens": 123, "output_tokens": 45,
  161. "cache_creation_input_tokens": 0,
  162. "cache_read_input_tokens": 0}
  163. lam = ConversationLam(name="t", provider=FakeProvider(), system_prompt="s")
  164. # 复用 _execute_agent 内的同款查找逻辑(提取为本测试的等价实现以
  165. # 锁定语义:_think_ref 优先 → 自身 get_usage → .agent 递归)
  166. def find(t, depth=0):
  167. if t is None or depth > 5:
  168. return None
  169. ref = getattr(t, "_think_ref", None)
  170. if ref is not None and callable(getattr(ref, "get_usage", None)):
  171. return ref
  172. if callable(getattr(t, "get_usage", None)):
  173. return t
  174. return find(getattr(t, "agent", None), depth + 1)
  175. assert find(lam) is lam
  176. guarded = Guard(lam, validator=lambda x: True)
  177. assert find(guarded) is lam
  178. wrapped = Memory(Guard(lam, validator=lambda x: True))
  179. assert find(wrapped) is lam
  180. assert find(wrapped).get_usage()["input_tokens"] == 123
  181. def test_anthropic_provider_accumulates_usage():
  182. """AnthropicProvider 现在有 get_usage/reset_session(与 openai_compat 同形)。"""
  183. from lambdagent.providers.anthropic_provider import AnthropicProvider
  184. from lambdagent.providers.base import ProviderConfig
  185. p = AnthropicProvider(ProviderConfig(model="claude-sonnet-4-6"))
  186. u = p.get_usage()
  187. assert u == {"input_tokens": 0, "output_tokens": 0,
  188. "cache_creation_input_tokens": 0,
  189. "cache_read_input_tokens": 0}
  190. p._usage_input += 10; p._usage_output += 5
  191. assert p.get_usage()["input_tokens"] == 10
  192. p.reset_session()
  193. assert p.get_usage()["output_tokens"] == 0