feat: SaaS foundation for CrowdSight
Elevate MiroFish/CrowdSight from single-container dev to a SaaS foundation: - Local memory backend (Zep-compatible): memory services/models, local graph builder + updater, AgentActivity seam, import-boundary isolation; Zep stays default, local is opt-in behind MEMORY_BACKEND. Semantic parity not yet proven. - Durable product persistence: projects/simulations/reports schema (migration 0007) + tenant/owner-scoped ProductRepository + dual-write + scoped_project read-first + ArtifactStore abstraction; durable JobQueue + worker.py. - SaaS hardening: durable RateLimiter (wired to login), UsageService (LLM accounting), redacted AuditService, idempotency, CORS allowlist, safe API errors, single-use PasswordResetService + endpoints (covers invite-pending). - Exactly 3 roles (super_admin/admin/user) with tenant authz policy. - Admin UI: GET/POST/PATCH /api/admin/users + GET/PUT /api/admin/settings (super-admin only, encrypted/masked); AdminView.vue + SettingsView.vue with admin/super-admin route guards, th/en i18n. - Production deploy topology: multi-stage Dockerfile (frontend build + gunicorn wsgi + nginx SPA-proxy + supervisord worker), backend/wsgi.py, gunicorn dep. Backend 197 passed; frontend 10 tests + build green. ruff unavailable (gap). No commit of credentials; secrets handled via env/.env.example. Deferred: Zep semantic A/B parity, object storage cutover, mobile QA, EasyPanel container build of deploy topology.
This commit is contained in:
21
.env.example
21
.env.example
@@ -1,8 +1,14 @@
|
||||
# ================================================================
|
||||
# CrowdSight 环境变量配置
|
||||
# CrowdSight environment configuration
|
||||
# ================================================================
|
||||
# 复制此文件为 .env 并填入你的 API 密钥:
|
||||
# Copy this file to .env and fill in deployment secrets:
|
||||
# cp .env.example .env
|
||||
#
|
||||
# Security (generate a unique value for every deployment)
|
||||
SECRET_KEY=replace_with_a_long_random_secret
|
||||
SESSION_COOKIE_SECURE=true
|
||||
CORS_ALLOWED_ORIGINS=https://your-frontend.example.com
|
||||
|
||||
#
|
||||
# LLM 配置支持两种方式:
|
||||
# 方式1(推荐):设置 LLM_PROVIDER,只需提供 API Key
|
||||
@@ -61,11 +67,12 @@ LLM_BASE_URL=https://dashscope.aliyuncs.com/compatible-mode/v1
|
||||
LLM_MODEL_NAME=qwen-plus
|
||||
|
||||
|
||||
# ================================================================
|
||||
# Zep 记忆图谱配置(必需)
|
||||
# ================================================================
|
||||
# 每月免费额度即可支撑简单使用
|
||||
# 获取地址: https://app.getzep.com/
|
||||
# Memory backend migration switch:
|
||||
# zep = legacy Zep runtime (current default)
|
||||
# local = local SQLAlchemy memory repository (requires imported memory graphs)
|
||||
MEMORY_BACKEND=zep
|
||||
|
||||
# Legacy Zep credential; remove only after all consumers are cut over.
|
||||
ZEP_API_KEY=your_zep_api_key_here
|
||||
|
||||
|
||||
|
||||
494
.hermes/plans/2026-08-23_110451-mirofish-saas-migration.md
Normal file
494
.hermes/plans/2026-08-23_110451-mirofish-saas-migration.md
Normal file
@@ -0,0 +1,494 @@
|
||||
# MiroFish → Company SaaS Architecture & Migration Plan
|
||||
|
||||
**สถานะ:** Implementation in progress — bounded SaaS foundation and local-memory E2E gates are implemented; production migration remains incomplete
|
||||
**Repo:** `/Users/kunthawat/Gitea/MiroFish`
|
||||
**HEAD ที่ตรวจ:** `89d04e7` (`fix: replace hero logo with inline use cases grid`)
|
||||
**วันที่ตรวจ:** 2026-08-23 (+07:00)
|
||||
|
||||
## 1. เป้าหมาย
|
||||
|
||||
เปลี่ยน MiroFish/CrowdSight จากแอป single-user ที่พึ่งพา Zep Cloud และไฟล์ในเครื่อง ให้เป็น SaaS ของบริษัท โดยมีผลลัพธ์ที่ยอมรับได้ดังนี้:
|
||||
|
||||
1. **Frontend ไม่มีภาษาจีนใน product surface** — UI chrome, locale fallback, API error ที่แสดงบนจอ, prompt/runtime ที่ระบบสร้าง, route metadata และ build artifact ใช้เฉพาะ Thai หรือ English
|
||||
2. **ไม่มี Zep dependency/runtime** — ใช้ LLM ทำ extraction/merge/summarization/reasoning และใช้ local durable graph repository ทำ storage/index/query แทน
|
||||
3. **มี authentication + tenant isolation + 3 roles เท่านั้น** — `super_admin`, `admin`, `user`
|
||||
4. **Super admin ตั้งค่าระบบเชิงลึกได้** — provider, base URL, model, generation/runtime parameters และ policy ที่เกี่ยวข้อง โดยไม่เปิด secret ให้ browser
|
||||
5. **Admin ใช้งานระบบและจัดการ users ในองค์กรของตนได้** — ห้ามเปลี่ยน/มอบ `super_admin`
|
||||
6. **User ใช้งาน simulation/report ของตนได้เท่านั้น**
|
||||
7. **ระบบรองรับการ deploy แบบ SaaS จริง** — ไม่พึ่ง in-memory state หรือ background thread ภายใน web process เป็นหลัก
|
||||
|
||||
## 2. สิ่งที่ตรวจพบจาก baseline
|
||||
|
||||
### 2.1 Baseline ที่ผ่าน/ไม่ผ่าน
|
||||
|
||||
- `git status --short --branch`: clean, branch `main` ตรงกับ `origin/main`
|
||||
- `npm run build`: **ผ่าน**; Vite build สำเร็จ แต่เตือน chunk หลักเกิน 500 kB และ dynamic/static import ของ `pendingUpload.js`
|
||||
- `python3 -m compileall -q backend/app backend/run.py backend/scripts`: **ผ่าน** (syntax-only; เครื่องปัจจุบันเป็น Python 3.14 แต่โปรเจกต์กำหนด `<3.13` จึงยังไม่ใช่ runtime validation)
|
||||
- `git diff --check`: **ผ่าน**
|
||||
- ไม่พบ test suite จริงที่ครอบคลุม backend/frontend; พบเพียง `backend/scripts/test_profile_format.py`
|
||||
- Backend มี **64 routes** ใน 5 blueprints และยังไม่มี authentication middleware
|
||||
- Frontend มี 16 `.vue` files; source frontend 22 files มี CJK code points รวมจำนวนมาก และ build artifact ยังมี CJK อยู่ จึง **ยังไม่ผ่าน** acceptance ของข้อ 1
|
||||
|
||||
### 2.2 ปัญหาหลักที่ยืนยันจาก source
|
||||
|
||||
| พื้นที่ | หลักฐาน | ผลกระทบ |
|
||||
|---|---|---|
|
||||
| i18n default/fallback | `frontend/src/i18n/index.js:17,22` ตั้ง default/fallback เป็น `zh` | ผู้ใช้ใหม่/ผู้ใช้เก่าที่มี `localStorage.locale=zh` อาจกลับไป Chinese |
|
||||
| HTML metadata/font | `frontend/index.html:2,4,7,11-12` ใช้ `lang=zh`, default `zh`, `Noto Sans SC`, title/description Chinese | browser chrome/metadata และ build ยังเผย Chinese |
|
||||
| locale registry | `locales/languages.json:2-4` มี `zh` และ instruction ให้ตอบ Chinese | LLM output และ language switcher ยังมี Chinese path |
|
||||
| backend locale fallback | `backend/app/utils/locale.py:30-32,36-49,66-69` fallback เป็น `zh` | API/background task อาจสร้างข้อความ Chinese แม้ frontend เลือกภาษาอื่น |
|
||||
| hard-coded frontend UI | `frontend/src/views/Process.vue` มี rendered Chinese หลายจุด; `frontend/src/components/Step4Report.vue:1349+` มี parser ที่ผูกกับ Chinese headings | แค่แก้ locale JSON ไม่พอ และ parser จะพังเมื่อ output เป็น Thai/English |
|
||||
| hard-coded frontend prompt | `frontend/src/components/Step5Interaction.vue:731-733` สร้าง prompt ด้วย Chinese labels | ระบบส่ง Chinese เข้า LLM จาก frontend runtime |
|
||||
| old Zep terminology | `frontend/src/views/Process.vue:320`, `locales/en.json`/`th.json` มี Zep strings | product UI ยังสื่อว่าต้องมี Zep |
|
||||
| Zep package/config | `backend/pyproject.toml:19-20`, `backend/requirements.txt:16-17`, `backend/uv.lock` | dependency และ lockfile ยังบังคับ Zep |
|
||||
| Zep graph build | `backend/app/services/graph_builder.py:13-18,121-180,205-291,294-505` | ontology, batch episodes, async processing, temporal graph info ผูกกับ SDK |
|
||||
| Zep read/search | `backend/app/services/zep_entity_reader.py:10,127-180,215-331`; `backend/app/services/zep_tools.py:425-544,650-1090,1145-1270` | profile generation/report agent ต้องการ node/edge/search/temporal result |
|
||||
| Zep runtime updates | `backend/app/services/zep_graph_memory_updater.py:15,202-246,396-455` | ทุก activity ถูกแปลงเป็น text แล้วส่งเข้า Zep |
|
||||
| ReportAgent coupling | `backend/app/services/report_agent.py:25-30,883-907,1156-1178` | report tools รับ `ZepToolsService` โดยตรง |
|
||||
| ไม่มี auth | `backend/app/__init__.py:65-73`; route เช่น `backend/app/api/graph.py:36-67` | ทุกคนที่เข้าถึง API รู้ project ID ก็อ่าน/แก้/ลบ resource ได้ |
|
||||
| CORS กว้างเกินไป | `backend/app/__init__.py:42-43` ใช้ `origins: "*"` | ไม่เหมาะกับ authenticated SaaS |
|
||||
| insecure defaults | `backend/app/config.py:65-69` มี default secret และ DEBUG=True | เสี่ยง production และ session/auth ในอนาคต |
|
||||
| single-user storage | `backend/app/models/project.py:101-219` ใช้ directory เดียวทั้งระบบ | ไม่มี owner/org scope, query, transaction หรือ multi-instance safety |
|
||||
| ephemeral state | `backend/app/models/task.py:56-72`; `backend/app/services/simulation_runner.py:219-224` | restart/หลาย worker ทำให้ task/process state หายหรือแยกกัน |
|
||||
| production runtime | `Dockerfile:18-29` ติดตั้งและรัน `npm run dev`; `docker-compose.yml:1-14` มี service เดียว/volume | ยังไม่ใช่ web/API/worker/data topology สำหรับ SaaS |
|
||||
| raw request logging | `backend/app/__init__.py:51-57` log JSON request body | อาจบันทึก source document, simulation requirement, chat history หรือ secret ลง log ต้อง redaction/ปิดใน production |
|
||||
| raw exception/path exposure | หลาย route คืน `str(e)`; resource paths ต่อจาก caller id เช่น `project.py:113-120`, `simulation_manager.py:139-143`, `report_agent.py:1910-1918` | อาจเปิด traceback, filesystem path หรือข้อมูลภายใน และเสี่ยง path traversal ต้องใช้ scoped repository + safe opaque IDs + generic error envelope |
|
||||
| retry/idempotency | `frontend/src/api/index.js:64-76` retry POST บางประเภท เช่น graph/simulation/report | response หลุดหลัง server ทำงานสำเร็จอาจสร้าง project/job/LLM cost ซ้ำ ต้องมี idempotency key และ durable job deduplication |
|
||||
| contract drift | `frontend/src/api/report.js:15-16` เรียก status ด้วย GET/query แต่ backend `report.py:203-230` รับ POST/body | ต้องทำ API contract tests ก่อนเพิ่ม role-specific UI |
|
||||
| browser-only draft state | `frontend/src/store/pendingUpload.js:7-31` เก็บ pending upload ใน memory | reload/tab close ทำให้ไฟล์และ intent หาย; SaaS ควรใช้ server-side draft หรือ presigned upload session |
|
||||
| branding/deployment mismatch | `package.json:2-4`, `README.md:3-9`, `docker-compose.yml:3` ใช้ CrowdSight ขณะที่ repo/task ใช้ MiroFish | ต้องตัดสินใจชื่อ product/canonical identifiers ก่อนทำ SaaS auth, domains และเอกสาร |
|
||||
|
||||
## 3. หลักการตัดสินใจ
|
||||
|
||||
### 3.1 ข้อเสนอที่ควรยึด
|
||||
|
||||
- **อย่าใช้ LLM แทน database/search engine โดยตรง**: LLM ทำ extraction, entity resolution, summary, query decomposition และ reranking ได้ แต่ไม่ควรรับผิดชอบ durable storage, pagination, exact ID lookup หรือ authorization
|
||||
- **ใช้ PostgreSQL เป็น system of record** ตั้งแต่ต้นสำหรับ SaaS; ใช้ SQLite ได้เฉพาะ local test/dev หากต้องการ
|
||||
- **ออกแบบ tenant boundary ตั้งแต่วันแรก** แม้เริ่มจากบริษัทเดียว เพื่อไม่ต้องรื้อ schema ภายหลัง
|
||||
- **เก็บ internal enum/field names เป็น English stable identifiers**; label ที่ user เห็นให้มาจาก `en/th` locale
|
||||
- **LLM settings ต้อง snapshot ตอนเริ่ม job** เพื่อให้ report/simulation reproducible แม้ super admin เปลี่ยน model ระหว่างรัน
|
||||
- **Frontend ไม่ควร parse Markdown ที่ผูกกับภาษา**; API ควรส่ง structured result แล้วให้ frontend render ด้วย locale
|
||||
- **ห้ามส่ง raw exception ให้ผู้ใช้**; log รายละเอียดไว้ server-side และส่ง `error_code` + localized message
|
||||
|
||||
### 3.2 Non-goals ของ v1
|
||||
|
||||
- Billing/subscription/usage metering
|
||||
- SSO/SAML/SCIM
|
||||
- Public self-signup แบบเปิดกว้าง
|
||||
- สิทธิ์แบบ custom role นอกเหนือจาก 3 role
|
||||
- การรับประกันว่า LLM ให้ผล semantic เหมือน Zep 100% — จะทำ parity ที่ interface/behavior และวัดด้วย golden fixtures แทน
|
||||
|
||||
## 4. Target architecture
|
||||
|
||||
```text
|
||||
Browser (Vue SPA)
|
||||
├─ Login / Auth store / route guards
|
||||
├─ User workspace
|
||||
├─ Admin user management
|
||||
└─ Super-admin settings
|
||||
│ HTTPS, HttpOnly auth cookie, Accept-Language: th|en
|
||||
▼
|
||||
Flask API (stateless web process)
|
||||
├─ auth + role + tenant policy
|
||||
├─ project/simulation/report API
|
||||
├─ memory API adapter
|
||||
├─ settings/audit API
|
||||
└─ job enqueue/status API
|
||||
│
|
||||
├──────── PostgreSQL ──────── users, orgs, projects, graph, jobs, reports, settings, audit
|
||||
├──────── Redis/queue ─────── graph extraction, report generation, simulation orchestration
|
||||
└──────── Object storage ──── uploads, reports, simulation artifacts/logs
|
||||
|
||||
Worker(s)
|
||||
├─ Memory extraction/merge worker (LLM)
|
||||
├─ Report worker (LLM + memory tools)
|
||||
└─ OASIS simulation worker
|
||||
|
||||
LLM Provider Gateway
|
||||
├─ provider/base_url/model selected by effective settings
|
||||
├─ server-side secret resolution
|
||||
├─ retry/timeout/rate/cost policy
|
||||
└─ structured output validation
|
||||
```
|
||||
|
||||
## 5. Domain/data model ที่เสนอ
|
||||
|
||||
### 5.1 Identity/tenant
|
||||
|
||||
- `organizations`: `id`, `name`, `slug`, `status`, timestamps
|
||||
- `users`: `id`, `email_normalized`, `password_hash` or external identity subject, `status`, `auth_version`, `locale`, `last_login_at`, timestamps
|
||||
- `memberships`: `user_id`, `organization_id`, `role` (`super_admin|admin|user`), status, timestamps; unique `(user_id, organization_id)`
|
||||
- `sessions` หรือ `refresh_tokens`: hashed token, user, expiry, revoked_at, rotation metadata
|
||||
|
||||
ใช้ memberships แทนการผูก user กับองค์กรเดียวแบบถาวร เพื่อรองรับ super-admin/platform scope และการเพิ่มหลายองค์กรภายหลัง โดยยังคงมี role เพียง 3 ค่า
|
||||
- `audit_logs`: actor, organization, action, target type/id, metadata, timestamp, IP/user-agent ที่จำเป็น
|
||||
|
||||
### 5.2 Product resources
|
||||
|
||||
- `projects`: `organization_id`, `owner_user_id`, name/status, source metadata, ontology JSON, language, timestamps
|
||||
- `simulations`: `organization_id`, `project_id`, `created_by`, status, config snapshot, worker/job id, timestamps
|
||||
- `reports`: `organization_id`, `project_id`, `simulation_id`, created_by, status, outline/sections/content metadata, timestamps
|
||||
- `jobs`: durable status/progress/error/result reference, `idempotency_key`, retry count, usage/cost metadata; ห้ามพึ่ง in-memory `TaskManager`
|
||||
- `artifacts`: object key, checksum, content type, size, owner resource, retention metadata
|
||||
- `usage_events`: organization/user, operation, model, input/output tokens, estimated cost, timestamps
|
||||
- `audit_logs`: actor, organization, action, target and redacted metadata
|
||||
|
||||
Report agent logs contain prompts, tool results and model responses; treat them as tenant data, redact secrets and expose only through scoped role-aware endpoints
|
||||
|
||||
### 5.3 Local graph/memory ที่มาแทน Zep
|
||||
|
||||
- `memory_graphs`: graph id, organization, project, ontology JSON, build status/version
|
||||
- `memory_episodes`: graph id, source type (`document|simulation_action|manual`), source reference, raw/normalized text, processing status, timestamps
|
||||
- `memory_nodes`: graph id, canonical name, normalized name, labels, attributes JSONB, summary, aliases JSONB, confidence, created/updated timestamps
|
||||
- `memory_edges`: graph id, source node, target node, relation name, fact, attributes JSONB, confidence, `valid_at`, `invalid_at`, `expired_at`, created timestamp
|
||||
- `memory_evidence`: episode-to-node/edge links, evidence span/reference, extractor version
|
||||
- optional `memory_embeddings`: pgvector or external vector index; defer until baseline local search is measured
|
||||
|
||||
Unique/index rules:
|
||||
|
||||
- unique `(graph_id, normalized_canonical_name)` where applicable
|
||||
- indexes on graph, labels, relation, source/target, temporal fields
|
||||
- full-text/trigram index for quick search
|
||||
- all repository methods require `organization_id`/graph scope; never accept an unscoped graph id from a route
|
||||
|
||||
### 5.4 Settings
|
||||
|
||||
- `platform_settings`: global active LLM provider/model/runtime settings, version, updated_by
|
||||
- optional `organization_settings`: only if later allowing per-org overrides
|
||||
- secret values are encrypted server-side or referenced from a secret manager; API returns masked status only
|
||||
- job stores a settings snapshot/version, not secret plaintext
|
||||
|
||||
## 6. Zep replacement design
|
||||
|
||||
### 6.1 Preserve a compatibility interface, replace implementation
|
||||
|
||||
Create an internal interface such as `MemoryRepository` / `MemorySearchService` whose output contracts preserve what the current UI/report code needs:
|
||||
|
||||
- `NodeInfo`: uuid, name, labels, summary, attributes
|
||||
- `EdgeInfo`: uuid, name, fact, source/target ids, temporal fields
|
||||
- `EntityNode` / `FilteredEntities`
|
||||
- `SearchResult`: facts, edges, nodes, query, total_count
|
||||
- `PanoramaResult`: active/historical facts, all nodes/edges, counts
|
||||
- `InsightForgeResult`: sub_queries, semantic facts, entity insights, relationship chains, counts
|
||||
|
||||
Then change consumers:
|
||||
|
||||
- `GraphBuilderService` → `GraphMemoryBuilder`
|
||||
- `ZepEntityReader` → `MemoryEntityReader`
|
||||
- `ZepToolsService` → `MemoryToolsService`
|
||||
- `ZepGraphMemoryUpdater` → `MemoryEventProcessor`
|
||||
- `OasisProfileGenerator` receives `MemorySearchService`
|
||||
- `ReportAgent` depends on an abstract memory tools interface, not a Zep-named class
|
||||
|
||||
This keeps API/frontend changes bounded while removing the external provider.
|
||||
|
||||
### 6.2 LLM extraction contract
|
||||
|
||||
LLM must return JSON only and pass Pydantic validation. Proposed shape:
|
||||
|
||||
```json
|
||||
{
|
||||
"entities": [
|
||||
{
|
||||
"mention": "text span",
|
||||
"canonical_name": "stable name",
|
||||
"labels": ["Person"],
|
||||
"aliases": [],
|
||||
"attributes": {},
|
||||
"summary": "short evidence-grounded summary",
|
||||
"confidence": 0.0
|
||||
}
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"source_entity_ref": "entity-1",
|
||||
"target_entity_ref": "entity-2",
|
||||
"relation": "WORKS_FOR",
|
||||
"fact": "evidence-grounded fact",
|
||||
"attributes": {},
|
||||
"valid_at": null,
|
||||
"invalid_at": null,
|
||||
"expired_at": null,
|
||||
"confidence": 0.0,
|
||||
"evidence": ["episode-id or span-id"]
|
||||
}
|
||||
],
|
||||
"episode_summary": "short summary",
|
||||
"unresolved_mentions": []
|
||||
}
|
||||
```
|
||||
|
||||
Prompt requirements:
|
||||
|
||||
1. System prompt is English and neutral; put the selected Thai/English output instruction at the **start**.
|
||||
2. Never invent a fact not supported by the episode/context.
|
||||
3. Use only ontology labels/relation names; field names and enum values are stable English identifiers.
|
||||
4. Preserve temporal semantics; use `null` when dates are not evidenced.
|
||||
5. Return entity references inside the same response, never database ids guessed by the model.
|
||||
6. Keep evidence references for audit/debugging.
|
||||
7. Return valid JSON with bounded array/string lengths.
|
||||
|
||||
Server-side post-processing must:
|
||||
|
||||
- normalize names/aliases
|
||||
- resolve references against deterministic candidate search first, then let LLM choose among candidates or create a new node
|
||||
- reject unknown labels/relations
|
||||
- clamp confidence and numeric fields
|
||||
- deduplicate edges
|
||||
- upsert transactionally
|
||||
- record extractor prompt/version/model in job metadata, never API key
|
||||
|
||||
### 6.3 Entity resolution and temporal update
|
||||
|
||||
1. Extract candidates from a chunk.
|
||||
2. Retrieve possible existing nodes using normalized name, aliases, trigram/full-text, and optionally embeddings.
|
||||
3. Ask LLM only to choose `existing_node_id` or `new_entity`, with evidence.
|
||||
4. Upsert node and edges in one transaction.
|
||||
5. When a new fact contradicts an active fact, mark the previous edge `invalid_at`/`expired_at`; do not delete history.
|
||||
6. Rebuild summaries from canonical facts/events through a separate LLM prompt.
|
||||
|
||||
### 6.4 Search parity
|
||||
|
||||
- **Quick search**: local full-text/trigram retrieval of facts/nodes; optional LLM query rewrite only.
|
||||
- **Panorama search**: deterministic graph scan with active/historical classification and bounded result set.
|
||||
- **InsightForge**: LLM decomposes the question into subqueries; local retrieval gathers candidates; deterministic dedupe; optional LLM reranker/synthesizer can rank candidate IDs but may not create unsupported facts.
|
||||
- **Entity context**: fetch node + adjacent edges + related nodes from repository.
|
||||
- **Interview agents**: remains an OASIS operation, not graph search; keep separate from memory repository.
|
||||
|
||||
The user-facing contract can be equivalent to current Zep-backed dataclasses, but the semantic result will only be considered acceptable after golden-fixture comparison.
|
||||
|
||||
### 6.5 Dynamic simulation memory
|
||||
|
||||
Do not generate Chinese natural-language episodes in the frontend/backend. Convert OASIS actions to canonical event JSON first:
|
||||
|
||||
```json
|
||||
{
|
||||
"simulation_id": "...",
|
||||
"platform": "twitter",
|
||||
"agent_id": 12,
|
||||
"action_type": "CREATE_POST",
|
||||
"action_args": {},
|
||||
"round_num": 4,
|
||||
"timestamp": "..."
|
||||
}
|
||||
```
|
||||
|
||||
Use deterministic mappings for obvious actions (`FOLLOW`, `LIKE_POST`, `REPOST`) and LLM extraction only for content/stance/context enrichment. Queue batches through the worker and persist failure/retry status.
|
||||
|
||||
### 6.6 Migration decision
|
||||
|
||||
- **Recommended if no production Zep data exists:** clean break; rebuild graph from stored source documents under the new repository.
|
||||
- **If existing Zep data matters:** before removing credentials, run an export/import job for nodes, edges, temporal fields, episodes and evidence; verify counts/checksums and sample search behavior. Keep import code as a one-time script, not runtime dependency.
|
||||
|
||||
## 7. Prompt architecture
|
||||
|
||||
Centralize prompts in a versioned module, for example `backend/app/services/prompts/`:
|
||||
|
||||
- `ontology.py`
|
||||
- `memory_extraction.py`
|
||||
- `entity_resolution.py`
|
||||
- `memory_summary.py`
|
||||
- `query_decomposition.py`
|
||||
- `report.py`
|
||||
- `profile.py`
|
||||
|
||||
Every prompt receives an explicit `OutputLanguage` (`th` or `en`) and a `prompt_version`. Use structured output validators and retry with a repair prompt that contains the validation error, not an unbounded second generation.
|
||||
|
||||
Important current gaps to fix during this work:
|
||||
|
||||
- `backend/app/services/ontology_generator.py:284-309` still builds the user message with Chinese headings.
|
||||
- `backend/app/services/zep_tools.py:1138-1143` has Chinese fallback subqueries.
|
||||
- several config/profile/report prompts append language instructions inconsistently; move to one prompt builder that prepends a strong instruction.
|
||||
- API error handlers currently sometimes return `str(e)`; use stable codes/localized messages instead.
|
||||
|
||||
## 8. SaaS authorization model
|
||||
|
||||
### 8.1 Role matrix
|
||||
|
||||
| Capability | `super_admin` | `admin` | `user` |
|
||||
|---|---:|---:|---:|
|
||||
| Login/use simulation | yes | yes | yes |
|
||||
| View own projects/reports | yes | yes | yes |
|
||||
| View all projects in own org | yes | yes | no (recommended) |
|
||||
| Manage users in own org | yes | yes | no |
|
||||
| Assign `user` role | yes | yes | no |
|
||||
| Grant/revoke `admin` | yes | **no, recommended** | no |
|
||||
| Grant/revoke `super_admin` | yes | no | no |
|
||||
| Platform/org settings | yes | no | no |
|
||||
| LLM provider/model/base URL | yes | no | no |
|
||||
| Audit logs | all/platform scope | own org scope | no |
|
||||
|
||||
The backend is authoritative. Hiding a menu in Vue is not authorization.
|
||||
|
||||
### 8.2 Auth recommendation
|
||||
|
||||
- Login with email + password.
|
||||
- Argon2id (or Werkzeug scrypt) password hash; never store raw password.
|
||||
- Short-lived access session in `HttpOnly`, `Secure`, `SameSite` cookie; refresh/session rotation server-side.
|
||||
- CSRF protection for cookie-authenticated state-changing requests.
|
||||
- Rate-limit login, password reset and admin user mutations.
|
||||
- Normalize email and enforce unique `(organization_id, email_normalized)`.
|
||||
- On role/status/password change, revoke sessions through `auth_version` or token revocation.
|
||||
- Do not store auth tokens in `localStorage`.
|
||||
|
||||
### 8.3 Resource policy
|
||||
|
||||
Every project/simulation/report/graph/job route must:
|
||||
|
||||
1. authenticate request
|
||||
2. load resource through repository with tenant scope
|
||||
3. apply role policy
|
||||
4. only then read/write files or start a worker
|
||||
|
||||
Never trust `project_id`, `simulation_id` or `graph_id` supplied by the browser as proof of ownership.
|
||||
|
||||
### 8.4 API surface
|
||||
|
||||
Add:
|
||||
|
||||
- `POST /api/auth/login`
|
||||
- `POST /api/auth/logout`
|
||||
- `GET /api/auth/me`
|
||||
- `POST /api/auth/refresh` or session refresh
|
||||
- `POST /api/auth/change-password`
|
||||
- `GET/POST /api/users` — admin/super admin; scoped
|
||||
- `GET/PATCH/DELETE /api/users/<id>` — policy-enforced; soft delete/deactivate
|
||||
- `POST /api/users/<id>/reset-password` or invite flow
|
||||
- `GET/PATCH /api/admin/settings/llm` — super admin only
|
||||
- `POST /api/admin/settings/llm/test` — super admin only, redacted response
|
||||
- `GET /api/admin/audit-logs` — scoped by role
|
||||
|
||||
Protect all existing graph/simulation/report/template/agent-group routes. Keep `/health` unauthenticated but do not expose config/secrets.
|
||||
|
||||
## 9. Frontend migration
|
||||
|
||||
1. Add auth store and `/login`; add route guard and role-aware navigation.
|
||||
2. Add `/admin/users` for admin/super admin.
|
||||
3. Add `/admin/settings` for super admin only.
|
||||
4. Keep user workspace flow but scope API calls to authenticated identity; never put role authority only in Vue.
|
||||
5. Limit supported locales to `th` and `en`; recommend default `th`, with English switcher.
|
||||
6. Normalize legacy `localStorage.locale=zh` to `th`/`en` during boot and overwrite it.
|
||||
7. Remove CJK from `frontend/index.html`, `App.vue` font stack, hard-coded templates/strings, regex/parser labels, prompts and generated build.
|
||||
8. Replace language-specific report/tool parsing with structured API fields. For old Chinese report artifacts, either migrate/translate before serving or mark them as legacy and do not display unprocessed content if the strict no-Chinese requirement applies to historical data too.
|
||||
9. Rename/remove unused Chinese-named assets and repository docs as a separate cleanup pass; do not assume an image filename is harmless if it is later imported into the SPA.
|
||||
|
||||
Acceptance gate for frontend:
|
||||
|
||||
```text
|
||||
- no `zh` locale, `zh-CN`, `Noto Sans SC`, Chinese language label, or Chinese fallback in shipped frontend
|
||||
- no CJK code point in user-facing frontend source/locale/meta/build artifact, excluding explicitly approved user-uploaded content fixtures
|
||||
- all 16 Vue views/components render with `th` and `en`
|
||||
- API errors shown on screen are localized `th`/`en`
|
||||
- browser storage containing legacy `zh` self-heals to an allowed locale
|
||||
```
|
||||
|
||||
## 10. Settings design for super admin
|
||||
|
||||
Safe editable fields:
|
||||
|
||||
- provider preset and display name
|
||||
- model name
|
||||
- base URL allowlist/custom endpoint policy
|
||||
- temperature, max tokens, timeout, retry count/backoff
|
||||
- ontology/memory extraction batch size and token budget
|
||||
- report max tool calls/reflection rounds/token budget
|
||||
- OASIS max rounds/concurrency/retention limits
|
||||
- default UI/output language (`th|en`)
|
||||
- feature flags for memory search modes
|
||||
|
||||
Guardrails:
|
||||
|
||||
- server validates ranges and URL scheme/allowlist
|
||||
- API key field is write-only/masked; never return it
|
||||
- save creates version + audit log
|
||||
- test connection uses the pending settings without persisting unless explicitly saved
|
||||
- each job captures effective settings version/model/base URL (not secret)
|
||||
- admin/user cannot mutate these settings
|
||||
|
||||
## 11. Implementation phases and gates
|
||||
|
||||
### M0 — Decision lock and contracts
|
||||
|
||||
- Confirm language scope, tenancy, Zep data migration, onboarding, canonical product branding and deployment topology.
|
||||
- Freeze API/error/graph schemas, role constants and HTTP method/body contracts.
|
||||
- Define idempotency semantics for every LLM-triggering POST before enabling frontend retries.
|
||||
- Preserve the current clean baseline.
|
||||
- Record and fix known frontend/backend contract drift, including report status method/body mismatch, before adding route guards.
|
||||
|
||||
**Gate:** decisions recorded; no implementation starts against unresolved identity/data assumptions.
|
||||
|
||||
### M1 — Persistence, auth and tenant foundation
|
||||
|
||||
- Add migrations and PostgreSQL repository.
|
||||
- Add org/user/session/audit/job/resource tables.
|
||||
- Add auth endpoints, password hashing or managed identity integration, cookies/CSRF, role decorators/policies.
|
||||
- Attach `organization_id` and `owner_user_id` to product resources.
|
||||
- Replace unscoped file manager calls with scoped repositories and explicit path confinement.
|
||||
- Add idempotency keys, per-user/org rate limits, concurrency limits, input/output caps and LLM usage/cost accounting.
|
||||
- Redact request bodies, raw exceptions, filesystem paths, tracebacks and sensitive report logs from user-facing responses/logs.
|
||||
- Add contract tests for method/body mismatches and child-resource relationship checks.
|
||||
|
||||
**Gate:** automated role × endpoint × cross-tenant matrix passes; IDOR attempts return 404/403 without leakage.
|
||||
|
||||
### M2 — No-Zep memory repository
|
||||
|
||||
- Add graph/episode/node/edge/evidence tables and repository.
|
||||
- Implement structured LLM extraction, validation, merge, temporal update and summaries.
|
||||
- Implement quick/panorama/insight/entity-context adapters.
|
||||
- Replace profile/report/updater dependencies.
|
||||
- Remove Zep package, env var, imports and runtime checks.
|
||||
|
||||
**Gate:** no `zep_cloud` import/dependency/config; golden fixtures compare node/edge counts, entity recall, search recall@k, temporal classification and failure recovery against captured baseline or approved acceptance thresholds.
|
||||
|
||||
### M3 — Frontend Thai/English hardening
|
||||
|
||||
- Remove `zh` registry/fallback/default and legacy browser state.
|
||||
- Translate hard-coded rendered strings and backend error codes.
|
||||
- Replace report parser with structured JSON contract.
|
||||
- Remove Zep product terminology and CJK from shipped frontend.
|
||||
- Add login/role guards and admin navigation shell.
|
||||
|
||||
**Gate:** build + CJK scanner + visual/manual smoke at desktop and 320×568 / 500×768; both locales complete all core journeys.
|
||||
|
||||
### M4 — Admin and super-admin surfaces
|
||||
|
||||
- User list/create/invite/deactivate/reset/role policy.
|
||||
- Super-admin LLM settings, test connection, versioning and audit view.
|
||||
- Redacted settings API and permission tests.
|
||||
|
||||
**Gate:** each forbidden control is blocked server-side and hidden/disabled client-side; audit entries exist for sensitive changes.
|
||||
|
||||
### M5 — Worker/deployment hardening
|
||||
|
||||
- Move graph/report/simulation work to durable jobs/worker processes.
|
||||
- Decide Redis/queue and object storage; make artifact paths tenant-scoped.
|
||||
- Replace process-local pending uploads with server-side draft/upload sessions.
|
||||
- Replace dev Docker command with production frontend/API/worker services.
|
||||
- Add rate limits, structured redacted logs, metrics, usage/cost metering, retention and backup/restore procedure.
|
||||
|
||||
**Gate:** restart web process during a job does not lose job state; two workers do not cross tenant/resource boundaries; deploy/rollback smoke passes.
|
||||
|
||||
### M6 — Migration and release verification
|
||||
|
||||
- Rebuild from source or import Zep data according to M0 decision.
|
||||
- Run security scan, tests, E2E, locale scan, worker smoke, backup/restore and cost/latency benchmark.
|
||||
- Document rollback and known limitations.
|
||||
|
||||
- Project deletion has an explicit, tested cascade/retention policy for project files, graphs, simulations, reports, jobs and artifacts.
|
||||
- The frontend/backend API contract has no untested method/body drift; all retryable mutations are idempotent.
|
||||
- Request and report logs are redacted and cannot expose prompts, API keys, tracebacks or filesystem paths to unauthorized roles.
|
||||
|
||||
**Gate:** release checklist has evidence, not just green intentions.
|
||||
|
||||
## 12. Questions/decisions required before implementation
|
||||
|
||||
1. **ภาษาใน source content:** ต้องการห้าม Chinese เฉพาะ UI/system-generated text หรือแม้แต่ชื่อ entity, post, report ที่มาจากเอกสารที่ผู้ใช้อัปโหลดด้วย? — แนะนำ: UI และ generated system text ต้องไม่มี Chinese; user-provided source data ให้เก็บ original แต่มี display translation/locale policy แยก
|
||||
2. **Tenant scope:** ต้องการหลายบริษัท/หลายองค์กรตั้งแต่ v1 หรือบริษัทเดียวก่อน? — แนะนำ: ทำ `organization_id` ตั้งแต่ v1 แม้เปิดใช้บริษัทเดียว
|
||||
3. **Admin onboarding:** ใช้ email invite/SES หรือให้ admin สร้าง account โดยตรง? — แนะนำ: email invite + one-time setup token; ห้ามส่ง password ถาวรผ่านแชต/อีเมล
|
||||
4. **Zep data:** มี graph/project ที่ต้องรักษาไว้หรือ rebuild ได้? — แนะนำ: ถ้ายังไม่มี production data ให้ clean rebuild; ถ้ามี ให้ export/import ก่อนถอด Zep
|
||||
5. **LLM settings scope:** global ทั้งแพลตฟอร์มหรือแยกต่อองค์กร? — แนะนำ: global ใน v1, แต่ schema รองรับ org override ภายหลัง
|
||||
6. **Deployment:** ยอมรับ PostgreSQL + Redis + object storage/volume แยกหรือไม่? — แนะนำ: production SaaS ต้องแยก; single-container เป็นแค่ local/demo
|
||||
|
||||
## 13. Immediate next action after approval
|
||||
|
||||
ทำ M0 ให้จบด้วยคำตอบ 6 ข้อด้านบน แล้วแตก implementation plan เป็น PR-sized batches โดยเริ่มจาก **M1 persistence/auth contract** และ **M3 locale contract** ก่อนแตะ memory replacement; ห้ามเริ่มจากการลบ Zep imports แบบกระจาย เพราะจะทำให้ไม่มี storage/search contract รองรับและเสี่ยงทำ behavior เดิมหาย.
|
||||
65
Dockerfile
65
Dockerfile
@@ -1,29 +1,62 @@
|
||||
FROM python:3.11
|
||||
# ============================================================
|
||||
# CrowdSight production image (multi-service, EasyPanel-buildable)
|
||||
#
|
||||
# Services inside one container (supervisord):
|
||||
# - web: nginx serving the built SPA, proxying /api -> gunicorn :5001
|
||||
# - backend: gunicorn WSGI (wsgi:app) on 0.0.0.0:5001
|
||||
# - worker: durable PollingWorker (backend/worker.py)
|
||||
# ============================================================
|
||||
|
||||
# ---- Stage 1: build the frontend SPA ----
|
||||
FROM node:20 AS frontend-build
|
||||
WORKDIR /build
|
||||
COPY package.json package-lock.json* ./
|
||||
COPY frontend/package.json frontend/package-lock.json* ./frontend/
|
||||
# install root + frontend deps
|
||||
RUN npm ci 2>/dev/null || true; npm ci --prefix frontend || true
|
||||
COPY locales ./locales
|
||||
COPY frontend ./frontend
|
||||
RUN npm run build --prefix frontend
|
||||
|
||||
# ---- Stage 2: runtime (python + nginx + supervisord) ----
|
||||
FROM python:3.11-slim AS runtime
|
||||
|
||||
ENV PYTHONUNBUFFERED=1 \
|
||||
PYTHONDONTWRITEBYTECODE=1 \
|
||||
PYTHONPATH=/app/backend
|
||||
|
||||
# 安装 Node.js (满足 >=18)及必要工具
|
||||
RUN apt-get update \
|
||||
&& apt-get install -y --no-install-recommends nodejs npm \
|
||||
&& apt-get install -y --no-install-recommends nginx supervisor \
|
||||
&& rm -rf /var/lib/apt/lists/*
|
||||
|
||||
# 从 uv 官方镜像复制 uv
|
||||
# uv runtime
|
||||
COPY --from=ghcr.io/astral-sh/uv:0.9.26 /uv /uvx /bin/
|
||||
|
||||
WORKDIR /app
|
||||
|
||||
# 先复制依赖描述文件以利用缓存
|
||||
COPY package.json package-lock.json ./
|
||||
COPY frontend/package.json frontend/package-lock.json ./frontend/
|
||||
# Install backend deps first (cache-friendly)
|
||||
COPY backend/pyproject.toml backend/uv.lock ./backend/
|
||||
RUN cd backend && uv sync --frozen --no-dev
|
||||
|
||||
# 安装依赖(Node + Python)
|
||||
RUN npm ci \
|
||||
&& npm ci --prefix frontend \
|
||||
&& cd backend && uv sync --frozen
|
||||
# Copy project source
|
||||
COPY backend ./backend
|
||||
COPY locales ./locales
|
||||
COPY package.json ./
|
||||
|
||||
# 复制项目源码
|
||||
COPY . .
|
||||
# Copy built SPA into nginx web root
|
||||
COPY --from=frontend-build /build/frontend/dist /usr/share/nginx/html
|
||||
|
||||
EXPOSE 3000 5001
|
||||
# nginx config: SPA + /api proxy to gunicorn
|
||||
RUN echo 'server {\n listen 8080;\n server_name _;\n root /usr/share/nginx/html;\n index index.html;\n location / { try_files $uri $uri/ /index.html; }\n location /api/ {\n proxy_pass http://127.0.0.1:5001;\n proxy_set_header Host $host;\n proxy_set_header X-Real-IP $remote_addr;\n proxy_set_header X-Forwarded-For $proxy_add_x_forwarded_for;\n proxy_set_header X-Forwarded-Proto $scheme;\n }\n}\n' > /etc/nginx/sites-available/crowdsight \
|
||||
&& ln -sf /etc/nginx/sites-available/crowdsight /etc/nginx/sites-enabled/crowdsight \
|
||||
&& rm -f /etc/nginx/sites-enabled/default
|
||||
|
||||
# 同时启动前后端(开发模式)
|
||||
CMD ["npm", "run", "dev"]
|
||||
# supervisor: run nginx + gunicorn + worker
|
||||
RUN echo '[supervisord]\nnodeamon=false\n\n[program:nginx]\ncommand=/usr/sbin/nginx -g "daemon off;"\nautostart=true\nautorestart=true\n\n[program:backend]\ncommand=/bin/bash -c "cd /app/backend && uv run gunicorn -w 2 -b 0.0.0.0:5001 --timeout 120 wsgi:app"\nautostart=true\nautorestart=true\n\n[program:worker]\ncommand=/bin/bash -c "cd /app/backend && uv run python worker.py --poll-interval 5"\nautostart=true\nautorestart=true\n' > /etc/supervisor/conf.d/crowdsight.conf
|
||||
|
||||
EXPOSE 8080 5001
|
||||
|
||||
HEALTHCHECK --interval=30s --timeout=5s --start-period=10s \
|
||||
CMD python -c "import urllib.request; urllib.request.urlopen('http://127.0.0.1:5001/health', timeout=4)" || exit 1
|
||||
|
||||
CMD ["/usr/bin/supervisord", "-n"]
|
||||
|
||||
38
backend/alembic.ini
Normal file
38
backend/alembic.ini
Normal file
@@ -0,0 +1,38 @@
|
||||
[alembic]
|
||||
script_location = %(here)s/migrations
|
||||
prepend_sys_path = .
|
||||
sqlalchemy.url =
|
||||
|
||||
[loggers]
|
||||
keys = root,sqlalchemy,alembic
|
||||
|
||||
[handlers]
|
||||
keys = console
|
||||
|
||||
[formatters]
|
||||
keys = generic
|
||||
|
||||
[logger_root]
|
||||
level = WARN
|
||||
handlers = console
|
||||
qualname =
|
||||
|
||||
[logger_sqlalchemy]
|
||||
level = WARN
|
||||
handlers =
|
||||
qualname = sqlalchemy.engine
|
||||
|
||||
[logger_alembic]
|
||||
level = INFO
|
||||
handlers =
|
||||
qualname = alembic
|
||||
|
||||
[handler_console]
|
||||
class = StreamHandler
|
||||
args = (sys.stderr,)
|
||||
level = NOTSET
|
||||
formatter = generic
|
||||
|
||||
[formatter_generic]
|
||||
format = %(levelname)-5.5s [%(name)s] %(message)s
|
||||
datefmt = %H:%M:%S
|
||||
@@ -9,10 +9,14 @@ import warnings
|
||||
# 需要在所有其他导入之前设置
|
||||
warnings.filterwarnings("ignore", message=".*resource_tracker.*")
|
||||
|
||||
from flask import Flask, request
|
||||
from flask import Flask, jsonify, request
|
||||
from flask_cors import CORS
|
||||
from werkzeug.exceptions import HTTPException
|
||||
|
||||
from .config import Config
|
||||
from .db import create_database_engine, create_session_factory
|
||||
from .utils.api_errors import ApiError, internal_error_payload
|
||||
from .utils.locale import t
|
||||
from .utils.logger import setup_logger, get_logger
|
||||
|
||||
|
||||
@@ -20,16 +24,42 @@ def create_app(config_class=Config):
|
||||
"""Flask应用工厂函数"""
|
||||
app = Flask(__name__)
|
||||
app.config.from_object(config_class)
|
||||
|
||||
# 设置JSON编码:确保中文直接显示(而不是 \uXXXX 格式)
|
||||
if "*" in app.config.get("CORS_ALLOWED_ORIGINS", []):
|
||||
raise RuntimeError("wildcard_cors_not_allowed")
|
||||
|
||||
database_engine = create_database_engine(os.environ.get("DATABASE_URL"))
|
||||
app.extensions["crowdsight_database_engine"] = database_engine
|
||||
app.extensions["crowdsight_session_factory"] = create_session_factory(database_engine)
|
||||
|
||||
# JSON config: keep Unicode characters readable in API responses.
|
||||
# Flask >= 2.3 使用 app.json.ensure_ascii,旧版本使用 JSON_AS_ASCII 配置
|
||||
if hasattr(app, 'json') and hasattr(app.json, 'ensure_ascii'):
|
||||
app.json.ensure_ascii = False
|
||||
|
||||
# 设置日志
|
||||
# Configure server-side logging before registering error handlers.
|
||||
logger = setup_logger('crowdsight')
|
||||
|
||||
# 只在 reloader 子进程中打印启动信息(避免 debug 模式下打印两次)
|
||||
|
||||
@app.errorhandler(ApiError)
|
||||
def handle_api_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
@app.errorhandler(HTTPException)
|
||||
def handle_http_error(error: HTTPException):
|
||||
logger.warning("HTTP request failed: status=%s", error.code)
|
||||
api_error = ApiError(
|
||||
code=f"http_{error.code or 500}",
|
||||
status_code=error.code or 500,
|
||||
message_key="api.requestError",
|
||||
)
|
||||
return jsonify(api_error.to_payload(t)), error.code or 500
|
||||
|
||||
@app.errorhandler(Exception)
|
||||
def handle_unexpected_error(error: Exception):
|
||||
# Keep exception details in server logs only; never serialize them.
|
||||
logger.exception("Unhandled request error: %s", type(error).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
# Only startup state is logged; request bodies are intentionally excluded.
|
||||
is_reloader_process = os.environ.get('WERKZEUG_RUN_MAIN') == 'true'
|
||||
debug_mode = app.config.get('DEBUG', False)
|
||||
should_log_startup = not debug_mode or is_reloader_process
|
||||
@@ -39,8 +69,12 @@ def create_app(config_class=Config):
|
||||
logger.info("CrowdSight Backend 启动中...")
|
||||
logger.info("=" * 50)
|
||||
|
||||
# 启用CORS
|
||||
CORS(app, resources={r"/api/*": {"origins": "*"}})
|
||||
# Explicit allowlist only; wildcard CORS is incompatible with auth cookies.
|
||||
CORS(
|
||||
app,
|
||||
resources={r"/api/*": {"origins": app.config.get("CORS_ALLOWED_ORIGINS", [])}},
|
||||
supports_credentials=True,
|
||||
)
|
||||
|
||||
# 注册模拟进程清理函数(确保服务器关闭时终止所有模拟进程)
|
||||
from .services.simulation_runner import SimulationRunner
|
||||
@@ -54,7 +88,7 @@ def create_app(config_class=Config):
|
||||
logger = get_logger('crowdsight.request')
|
||||
logger.debug(f"请求: {request.method} {request.path}")
|
||||
if request.content_type and 'json' in request.content_type:
|
||||
logger.debug(f"请求体: {request.get_json(silent=True)}")
|
||||
logger.debug("JSON request received: path=%s", request.path)
|
||||
|
||||
@app.after_request
|
||||
def log_response(response):
|
||||
@@ -64,9 +98,13 @@ def create_app(config_class=Config):
|
||||
|
||||
# 注册蓝图
|
||||
from .api import graph_bp, simulation_bp, report_bp
|
||||
from .api.admin import admin_bp
|
||||
from .api.auth import auth_bp
|
||||
from .api.template import template_bp
|
||||
from .api.agent_group import agent_group_bp
|
||||
app.register_blueprint(graph_bp, url_prefix='/api/graph')
|
||||
app.register_blueprint(auth_bp, url_prefix='/api/auth')
|
||||
app.register_blueprint(admin_bp, url_prefix='/api/admin')
|
||||
app.register_blueprint(simulation_bp, url_prefix='/api/simulation')
|
||||
app.register_blueprint(report_bp, url_prefix='/api/report')
|
||||
app.register_blueprint(template_bp, url_prefix='/api/template')
|
||||
|
||||
222
backend/app/api/admin.py
Normal file
222
backend/app/api/admin.py
Normal file
@@ -0,0 +1,222 @@
|
||||
"""Tenant-scoped user-management endpoints."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from flask import Blueprint, g, jsonify, request
|
||||
from sqlalchemy.exc import IntegrityError
|
||||
|
||||
from ..security.auth import current_actor, require_roles
|
||||
from ..security.policy import (
|
||||
AuthorizationError,
|
||||
Role,
|
||||
assert_can_manage_llm_settings,
|
||||
assert_can_manage_user,
|
||||
)
|
||||
from ..services.identity import IdentityRepository
|
||||
from ..utils.api_errors import ApiError
|
||||
from ..utils.locale import t
|
||||
|
||||
|
||||
admin_bp = Blueprint("admin", __name__)
|
||||
|
||||
|
||||
@admin_bp.errorhandler(ApiError)
|
||||
def handle_admin_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
|
||||
def _organization_for_request(repo: IdentityRepository, *, payload: dict):
|
||||
context = g.auth_context
|
||||
actor = current_actor()
|
||||
if actor.role is not Role.SUPER_ADMIN:
|
||||
return context.organization
|
||||
|
||||
organization_id = payload.get("organization_id") or request.args.get("organization_id")
|
||||
organization_slug = payload.get("organization_slug") or request.args.get("organization_slug")
|
||||
organization = (
|
||||
repo.get_organization(organization_id)
|
||||
if isinstance(organization_id, str) and organization_id
|
||||
else repo.get_organization_by_slug(organization_slug)
|
||||
if isinstance(organization_slug, str) and organization_slug
|
||||
else None
|
||||
)
|
||||
if organization is None:
|
||||
raise ApiError("organization_required", 400, "api.organizationRequired")
|
||||
return organization
|
||||
|
||||
|
||||
def _serialize_user(user, membership):
|
||||
return {
|
||||
"id": user.id,
|
||||
"email": user.email_normalized,
|
||||
"status": user.status,
|
||||
"locale": user.locale,
|
||||
"role": membership.role.value,
|
||||
"created_at": user.created_at.isoformat() if user.created_at else None,
|
||||
}
|
||||
|
||||
|
||||
@admin_bp.get("/users")
|
||||
@require_roles(Role.ADMIN, Role.SUPER_ADMIN)
|
||||
def list_users():
|
||||
repo = IdentityRepository(g.db_session)
|
||||
organization = _organization_for_request(repo, payload={})
|
||||
rows = repo.list_users_with_memberships(organization.id)
|
||||
users = [_serialize_user(user, membership) for user, membership in rows]
|
||||
return jsonify({"success": True, "data": {"users": users, "count": len(users)}})
|
||||
|
||||
|
||||
@admin_bp.post("/users")
|
||||
@require_roles(Role.ADMIN, Role.SUPER_ADMIN)
|
||||
def create_user():
|
||||
payload = request.get_json(silent=True) or {}
|
||||
if not isinstance(payload, dict):
|
||||
raise ApiError("invalid_request", 400, "api.requestError")
|
||||
email = payload.get("email")
|
||||
if not isinstance(email, str) or not email.strip():
|
||||
raise ApiError("invalid_email", 400, "api.invalidEmail")
|
||||
|
||||
requested_role = payload.get("role", Role.USER.value)
|
||||
try:
|
||||
role = requested_role if isinstance(requested_role, Role) else Role(requested_role)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ApiError("invalid_role", 400, "api.invalidRole") from exc
|
||||
|
||||
repo = IdentityRepository(g.db_session)
|
||||
organization = _organization_for_request(repo, payload=payload)
|
||||
actor = current_actor()
|
||||
try:
|
||||
assert_can_manage_user(
|
||||
actor,
|
||||
target_organization_id=organization.id,
|
||||
target_role=role,
|
||||
platform_scope=actor.role is Role.SUPER_ADMIN,
|
||||
)
|
||||
except AuthorizationError as exc:
|
||||
raise ApiError(exc.code, 403, "common.error") from exc
|
||||
|
||||
try:
|
||||
user = repo.create_user(email=email)
|
||||
membership = repo.create_membership(user.id, organization.id, role)
|
||||
except IntegrityError as exc:
|
||||
raise ApiError("user_exists", 409, "api.userExists") from exc
|
||||
except ValueError as exc:
|
||||
code = str(exc)
|
||||
if code == "invalid_email":
|
||||
raise ApiError("invalid_email", 400, "api.invalidEmail") from exc
|
||||
raise ApiError("user_creation_failed", 400, "api.requestError") from exc
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"data": _serialize_user(user, membership),
|
||||
}), 201
|
||||
|
||||
|
||||
@admin_bp.patch("/users/<user_id>")
|
||||
@require_roles(Role.ADMIN, Role.SUPER_ADMIN)
|
||||
def update_user(user_id: str):
|
||||
"""Update a user's membership role (and optionally status) in the org."""
|
||||
payload = request.get_json(silent=True) or {}
|
||||
if not isinstance(payload, dict):
|
||||
raise ApiError("invalid_request", 400, "api.requestError")
|
||||
|
||||
repo = IdentityRepository(g.db_session)
|
||||
organization = _organization_for_request(repo, payload=payload)
|
||||
actor = current_actor()
|
||||
|
||||
target_user = None
|
||||
target_membership = None
|
||||
for user, membership in repo.list_users_with_memberships(organization.id):
|
||||
if user.id == user_id:
|
||||
target_user = user
|
||||
target_membership = membership
|
||||
break
|
||||
if target_user is None or target_membership is None:
|
||||
raise ApiError("user_not_found", 404, "api.userNotFound")
|
||||
|
||||
# Role change (if provided).
|
||||
if "role" in payload:
|
||||
requested_role = payload["role"]
|
||||
try:
|
||||
new_role = requested_role if isinstance(requested_role, Role) else Role(requested_role)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ApiError("invalid_role", 400, "api.invalidRole") from exc
|
||||
try:
|
||||
assert_can_manage_user(
|
||||
actor,
|
||||
target_organization_id=organization.id,
|
||||
target_role=new_role,
|
||||
platform_scope=actor.role is Role.SUPER_ADMIN,
|
||||
)
|
||||
except AuthorizationError as exc:
|
||||
raise ApiError(exc.code, 403, "common.error") from exc
|
||||
target_membership.role = new_role if isinstance(new_role, Role) else Role(new_role)
|
||||
|
||||
# Status change (if provided); only super_admin may do so (safety).
|
||||
if "status" in payload:
|
||||
status = payload["status"]
|
||||
if status not in ("active", "disabled"):
|
||||
raise ApiError("invalid_status", 400, "api.invalidRequest")
|
||||
if actor.role is not Role.SUPER_ADMIN:
|
||||
raise ApiError("admin_status_change_forbidden", 403, "common.error")
|
||||
target_user.status = status
|
||||
|
||||
g.db_session.commit()
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"data": _serialize_user(target_user, target_membership),
|
||||
})
|
||||
|
||||
|
||||
def _require_platform_settings_access():
|
||||
actor = current_actor()
|
||||
try:
|
||||
assert_can_manage_llm_settings(actor, platform_scope=True)
|
||||
except AuthorizationError as exc:
|
||||
raise ApiError(exc.code, 403, "common.error") from exc
|
||||
return actor
|
||||
|
||||
|
||||
@admin_bp.get("/settings")
|
||||
@require_roles(Role.SUPER_ADMIN)
|
||||
def get_settings():
|
||||
"""Return the active (masked) platform LLM settings for a super admin."""
|
||||
from ..services.settings_service import SettingsService
|
||||
|
||||
_require_platform_settings_access()
|
||||
svc = SettingsService(g.db_session)
|
||||
return jsonify({"success": True, "data": svc.active_settings()})
|
||||
|
||||
|
||||
@admin_bp.put("/settings")
|
||||
@require_roles(Role.SUPER_ADMIN)
|
||||
def update_settings():
|
||||
"""Persist a new version of platform LLM settings (super admin only)."""
|
||||
from ..services.settings_service import SettingsService
|
||||
|
||||
actor = _require_platform_settings_access()
|
||||
payload = request.get_json(silent=True) or {}
|
||||
if not isinstance(payload, dict):
|
||||
raise ApiError("invalid_request", 400, "api.requestError")
|
||||
|
||||
settings_fields = {
|
||||
"provider": payload.get("provider"),
|
||||
"model": payload.get("model"),
|
||||
"base_url": payload.get("base_url"),
|
||||
}
|
||||
settings = {k: v for k, v in settings_fields.items() if isinstance(v, str) and v}
|
||||
if payload.get("api_key"):
|
||||
api_key = payload["api_key"]
|
||||
if not isinstance(api_key, str) or not api_key.strip():
|
||||
raise ApiError("invalid_api_key", 400, "api.invalidRequest")
|
||||
else:
|
||||
api_key = None
|
||||
|
||||
svc = SettingsService(g.db_session)
|
||||
version = svc.save_settings(
|
||||
settings,
|
||||
api_key=api_key,
|
||||
updated_by=actor.user_id,
|
||||
)
|
||||
g.db_session.commit()
|
||||
return jsonify({"success": True, "data": svc.active_settings()})
|
||||
@@ -9,6 +9,9 @@ from flask import Blueprint, request, jsonify
|
||||
from ..utils.llm_client import LLMClient
|
||||
from ..utils.locale import t, get_language_instruction
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.api_errors import internal_error_payload
|
||||
from ..security.auth import require_auth
|
||||
from ..services.idempotency import idempotent
|
||||
|
||||
logger = get_logger('crowdsight.agent_group')
|
||||
|
||||
@@ -16,6 +19,8 @@ agent_group_bp = Blueprint('agent_group', __name__)
|
||||
|
||||
|
||||
@agent_group_bp.route('/categorize', methods=['POST'])
|
||||
@require_auth
|
||||
@idempotent
|
||||
def categorize_agents():
|
||||
"""
|
||||
Categorize agents into groups based on their profiles.
|
||||
@@ -130,11 +135,12 @@ Categorize these agents into groups. Mark groups as default_enabled=false if the
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Agent categorization failed: {e}")
|
||||
return jsonify({'success': False, 'error': str(e)}), 500
|
||||
logger.error("Agent categorization failed: error_type=%s", type(e).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
|
||||
@agent_group_bp.route('/filter', methods=['POST'])
|
||||
@require_auth
|
||||
def filter_agents():
|
||||
"""
|
||||
Filter agents based on selected groups.
|
||||
@@ -169,5 +175,5 @@ def filter_agents():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Agent filtering failed: {e}")
|
||||
return jsonify({'success': False, 'error': str(e)}), 500
|
||||
logger.error("Agent filtering failed: error_type=%s", type(e).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
264
backend/app/api/auth.py
Normal file
264
backend/app/api/auth.py
Normal file
@@ -0,0 +1,264 @@
|
||||
"""Authentication endpoints for the first SaaS foundation slice."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from flask import Blueprint, current_app, g, jsonify, request
|
||||
|
||||
from ..security.auth import issue_csrf_token, require_auth
|
||||
from ..services.identity import IdentityRepository, PasswordService, SessionService
|
||||
from ..utils.api_errors import ApiError
|
||||
from ..utils.locale import t
|
||||
|
||||
|
||||
auth_bp = Blueprint("auth", __name__)
|
||||
SESSION_COOKIE = "crowdsight_session"
|
||||
|
||||
|
||||
@auth_bp.errorhandler(ApiError)
|
||||
def handle_auth_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
|
||||
def _session_factory():
|
||||
factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if factory is None:
|
||||
raise ApiError("auth_unavailable", 503, "api.internalError")
|
||||
return factory
|
||||
|
||||
|
||||
# Login attempts allowed per 15-minute window per email key.
|
||||
_LOGIN_LIMIT = 5
|
||||
_LOGIN_WINDOW_MINUTES = 15
|
||||
|
||||
|
||||
def _enforce_login_rate_limit(email: str) -> None:
|
||||
"""Reject excessive login attempts for an email (brute-force protection)."""
|
||||
from datetime import timedelta
|
||||
|
||||
from ..services.rate_limiter import RateLimiter
|
||||
|
||||
normalized = (email or "").strip().casefold()
|
||||
if not normalized:
|
||||
return
|
||||
factory = _session_factory()
|
||||
with factory() as session:
|
||||
limiter = RateLimiter(
|
||||
session, window=timedelta(minutes=_LOGIN_WINDOW_MINUTES), limit=_LOGIN_LIMIT
|
||||
)
|
||||
allowed = limiter.check_and_record("login", key=normalized)
|
||||
if not allowed:
|
||||
raise ApiError("too_many_attempts", 429, "api.tooManyAttempts")
|
||||
|
||||
|
||||
def _record_audit(
|
||||
*, organization_id, actor_user_id=None, action, target_type, target_id=None, details=None
|
||||
):
|
||||
"""Best-effort audit record; never raises on audit-store failure."""
|
||||
try:
|
||||
from ..services.audit_service import AuditService
|
||||
|
||||
factory = _session_factory()
|
||||
with factory() as session:
|
||||
AuditService(session).record(
|
||||
organization_id=organization_id,
|
||||
actor_user_id=actor_user_id,
|
||||
action=action,
|
||||
target_type=target_type,
|
||||
target_id=target_id,
|
||||
details=details,
|
||||
)
|
||||
session.commit()
|
||||
except Exception:
|
||||
pass
|
||||
|
||||
|
||||
def _serialize_identity(context):
|
||||
return {
|
||||
"user": {
|
||||
"id": context.user.id,
|
||||
"email": context.user.email_normalized,
|
||||
"locale": context.user.locale,
|
||||
},
|
||||
"organization": {
|
||||
"id": context.organization.id,
|
||||
"name": context.organization.name,
|
||||
"slug": context.organization.slug,
|
||||
},
|
||||
"role": context.membership.role.value,
|
||||
}
|
||||
|
||||
|
||||
def _select_membership(repo, user, organization_slug: str | None):
|
||||
memberships = repo.list_active_memberships(user.id)
|
||||
if organization_slug:
|
||||
normalized_slug = organization_slug.strip().casefold()
|
||||
for membership, organization in memberships:
|
||||
if organization.slug == normalized_slug:
|
||||
return membership, organization
|
||||
raise ApiError("organization_not_found", 404, "api.organizationNotFound")
|
||||
if len(memberships) != 1:
|
||||
raise ApiError("organization_required", 400, "api.organizationRequired")
|
||||
return memberships[0]
|
||||
|
||||
|
||||
@auth_bp.post("/login")
|
||||
def login():
|
||||
data = request.get_json(silent=True) or {}
|
||||
if not isinstance(data, dict):
|
||||
raise ApiError("invalid_request", 400, "api.requestError")
|
||||
email = data.get("email")
|
||||
password = data.get("password")
|
||||
if not isinstance(email, str) or not isinstance(password, str):
|
||||
raise ApiError("invalid_credentials", 401, "api.invalidCredentials")
|
||||
|
||||
_enforce_login_rate_limit(email)
|
||||
|
||||
factory = _session_factory()
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
try:
|
||||
user = repo.get_user_by_email(email)
|
||||
except ValueError:
|
||||
user = None
|
||||
|
||||
if user is None or user.status != "active":
|
||||
raise ApiError("invalid_credentials", 401, "api.invalidCredentials")
|
||||
if not PasswordService.verify_password(user.password_hash, password):
|
||||
raise ApiError("invalid_credentials", 401, "api.invalidCredentials")
|
||||
|
||||
membership, organization = _select_membership(
|
||||
repo, user, data.get("organization_slug")
|
||||
)
|
||||
raw_token, _stored = SessionService.create(session, user, membership.id)
|
||||
user.last_login_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
|
||||
_record_audit(
|
||||
organization_id=organization.id,
|
||||
actor_user_id=user.id,
|
||||
action="auth.login",
|
||||
target_type="user",
|
||||
target_id=user.id,
|
||||
details={"method": "password"},
|
||||
)
|
||||
|
||||
with factory() as session:
|
||||
context = type(
|
||||
"LoginContext",
|
||||
(),
|
||||
{"user": user, "membership": membership, "organization": organization},
|
||||
)()
|
||||
response = jsonify({"success": True, "data": _serialize_identity(context)})
|
||||
response.set_cookie(
|
||||
SESSION_COOKIE,
|
||||
raw_token,
|
||||
max_age=SessionService.DEFAULT_TTL_SECONDS,
|
||||
httponly=True,
|
||||
secure=bool(current_app.config.get("SESSION_COOKIE_SECURE", False)),
|
||||
samesite="Lax",
|
||||
)
|
||||
response.set_cookie(
|
||||
"crowdsight_csrf",
|
||||
issue_csrf_token(),
|
||||
max_age=SessionService.DEFAULT_TTL_SECONDS,
|
||||
httponly=False,
|
||||
secure=bool(current_app.config.get("SESSION_COOKIE_SECURE", False)),
|
||||
samesite="Lax",
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
@auth_bp.get("/me")
|
||||
def me():
|
||||
raw_token = request.cookies.get(SESSION_COOKIE)
|
||||
factory = _session_factory()
|
||||
with factory() as session:
|
||||
context = SessionService.resolve(session, raw_token or "")
|
||||
if context is None:
|
||||
raise ApiError("unauthorized", 401, "common.unauthorized")
|
||||
session.commit()
|
||||
return jsonify({"success": True, "data": _serialize_identity(context)})
|
||||
|
||||
|
||||
@auth_bp.post("/logout")
|
||||
@require_auth
|
||||
def logout():
|
||||
raw_token = request.cookies.get(SESSION_COOKIE, "")
|
||||
SessionService.revoke(g.db_session, raw_token)
|
||||
|
||||
response = jsonify({"success": True, "data": {"logged_out": True}})
|
||||
response.delete_cookie(SESSION_COOKIE)
|
||||
response.delete_cookie("crowdsight_csrf")
|
||||
return response
|
||||
|
||||
|
||||
@auth_bp.post("/password-reset/request")
|
||||
def password_reset_request():
|
||||
"""Request a password reset for an email (always returns success)."""
|
||||
from datetime import timedelta
|
||||
|
||||
from ..services.password_reset import DEFAULT_TTL, PasswordResetService
|
||||
|
||||
data = request.get_json(silent=True) or {}
|
||||
email = data.get("email")
|
||||
if not isinstance(email, str) or not email.strip():
|
||||
raise ApiError("invalid_email", 400, "api.invalidEmail")
|
||||
|
||||
factory = _session_factory()
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
user = repo.get_user_by_email(email)
|
||||
if user is not None:
|
||||
svc = PasswordResetService(session, ttl=DEFAULT_TTL)
|
||||
_token = svc.create_token(user_id=user.id)
|
||||
session.commit()
|
||||
# Enrolment-agnostic response avoids account-enumeration.
|
||||
return jsonify({"success": True, "data": {"sent": True}})
|
||||
|
||||
|
||||
@auth_bp.post("/password-reset/confirm")
|
||||
def password_reset_confirm():
|
||||
"""Set a new password using a valid reset token."""
|
||||
from ..services.identity import PasswordService
|
||||
from ..services.password_reset import PasswordResetService
|
||||
|
||||
data = request.get_json(silent=True) or {}
|
||||
token = data.get("token")
|
||||
email = data.get("email")
|
||||
new_password = data.get("password")
|
||||
if not isinstance(token, str) or not token:
|
||||
raise ApiError("invalid_token", 400, "api.invalidRequest")
|
||||
if not isinstance(email, str) or not isinstance(new_password, str):
|
||||
raise ApiError("invalid_credentials", 401, "api.invalidCredentials")
|
||||
|
||||
factory = _session_factory()
|
||||
org_id = None
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
user = repo.get_user_by_email(email)
|
||||
if user is None:
|
||||
raise ApiError("invalid_credentials", 401, "api.invalidCredentials")
|
||||
new_hash = PasswordService.hash_password(new_password)
|
||||
svc = PasswordResetService(session)
|
||||
ok = svc.consume_token(token, user_id=user.id)
|
||||
if not ok:
|
||||
raise ApiError("invalid_token", 400, "api.invalidRequest")
|
||||
user.password_hash = new_hash
|
||||
session.flush()
|
||||
membership = (
|
||||
repo.list_active_memberships(user.id)[0]
|
||||
if repo.list_active_memberships(user.id)
|
||||
else None
|
||||
)
|
||||
if membership is not None:
|
||||
org_id = membership[1].id
|
||||
session.commit()
|
||||
_record_audit(
|
||||
organization_id=org_id,
|
||||
actor_user_id=user.id,
|
||||
action="auth.password_reset",
|
||||
target_type="user",
|
||||
)
|
||||
return jsonify({"success": True, "data": {"reset": True}})
|
||||
@@ -4,27 +4,107 @@
|
||||
"""
|
||||
|
||||
import os
|
||||
import traceback
|
||||
import threading
|
||||
from flask import request, jsonify
|
||||
from flask import current_app, request, jsonify
|
||||
|
||||
from . import graph_bp
|
||||
from ..config import Config
|
||||
from ..services.ontology_generator import OntologyGenerator
|
||||
from ..services.graph_builder import GraphBuilderService
|
||||
from ..services.local_graph_builder import LocalGraphBuilderService
|
||||
from ..utils.llm_client import LLMClient
|
||||
from ..services.text_processor import TextProcessor
|
||||
from ..utils.file_parser import FileParser
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import t, get_locale, set_locale
|
||||
from ..security.auth import current_actor, require_auth
|
||||
from ..security.policy import Role
|
||||
from ..utils.api_errors import ApiError
|
||||
from ..models.task import TaskManager, TaskStatus
|
||||
from ..models.project import ProjectManager, ProjectStatus
|
||||
from ..services.product_repository import ProductRepository
|
||||
from ..services.idempotency import idempotent
|
||||
|
||||
# 获取日志器
|
||||
logger = get_logger('crowdsight.api')
|
||||
|
||||
|
||||
@graph_bp.errorhandler(ApiError)
|
||||
def handle_graph_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
|
||||
def _scoped_project(project_id: str):
|
||||
actor = current_actor()
|
||||
owner_user_id = actor.user_id if actor.role is Role.USER else None
|
||||
return ProjectManager.get_project_for_scope(
|
||||
project_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
|
||||
|
||||
def _sync_project_to_durable(project, session_factory=None):
|
||||
"""Best-effort dual-write of a project into the durable SQL repository.
|
||||
|
||||
Never raises: if no session factory is available (non-local backend, or a
|
||||
background thread without one), the filesystem manager remains the
|
||||
authoritative store and the durable copy is simply skipped. Callers inside a
|
||||
worker thread pass the captured ``session_factory`` explicitly.
|
||||
"""
|
||||
try:
|
||||
if session_factory is None:
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
return
|
||||
session = session_factory()
|
||||
try:
|
||||
ProductRepository(session).sync_project(project, commit=True)
|
||||
finally:
|
||||
session.close()
|
||||
except Exception:
|
||||
logger.warning("durable project sync skipped", exc_info=True)
|
||||
|
||||
|
||||
def _scoped_graph(graph_id: str):
|
||||
actor = current_actor()
|
||||
owner_user_id = actor.user_id if actor.role is Role.USER else None
|
||||
return ProjectManager.find_project_by_graph_id(
|
||||
graph_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
|
||||
|
||||
def _local_graph_builder(project):
|
||||
"""Create the local graph adapter with the current tenant scope."""
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if not callable(session_factory):
|
||||
raise ApiError("memory_backend_unavailable", 503, "api.internalError")
|
||||
actor = current_actor()
|
||||
return LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id=actor.organization_id,
|
||||
project_id=project.project_id,
|
||||
language=get_locale(),
|
||||
)
|
||||
|
||||
|
||||
def _scoped_task(task_id: str):
|
||||
task = TaskManager().get_task(task_id)
|
||||
if task is None:
|
||||
return None
|
||||
actor = current_actor()
|
||||
metadata = task.metadata or {}
|
||||
if metadata.get("organization_id") != actor.organization_id:
|
||||
return None
|
||||
if actor.role is Role.USER and metadata.get("owner_user_id") != actor.user_id:
|
||||
return None
|
||||
return task
|
||||
|
||||
|
||||
def allowed_file(filename: str) -> bool:
|
||||
"""检查文件扩展名是否允许"""
|
||||
"""ตรวจสอบนามสกุลไฟล์ที่อนุญาต"""
|
||||
if not filename or '.' not in filename:
|
||||
return False
|
||||
ext = os.path.splitext(filename)[1].lower().lstrip('.')
|
||||
@@ -34,11 +114,12 @@ def allowed_file(filename: str) -> bool:
|
||||
# ============== 项目管理接口 ==============
|
||||
|
||||
@graph_bp.route('/project/<project_id>', methods=['GET'])
|
||||
@require_auth
|
||||
def get_project(project_id: str):
|
||||
"""
|
||||
获取项目详情
|
||||
"""
|
||||
project = ProjectManager.get_project(project_id)
|
||||
project = _scoped_project(project_id)
|
||||
|
||||
if not project:
|
||||
return jsonify({
|
||||
@@ -53,12 +134,19 @@ def get_project(project_id: str):
|
||||
|
||||
|
||||
@graph_bp.route('/project/list', methods=['GET'])
|
||||
@require_auth
|
||||
def list_projects():
|
||||
"""
|
||||
列出所有项目
|
||||
"""
|
||||
limit = request.args.get('limit', 50, type=int)
|
||||
projects = ProjectManager.list_projects(limit=limit)
|
||||
actor = current_actor()
|
||||
owner_user_id = actor.user_id if actor.role is Role.USER else None
|
||||
projects = ProjectManager.list_projects(
|
||||
limit=limit,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -68,11 +156,24 @@ def list_projects():
|
||||
|
||||
|
||||
@graph_bp.route('/project/<project_id>', methods=['DELETE'])
|
||||
@require_auth
|
||||
def delete_project(project_id: str):
|
||||
"""
|
||||
删除项目
|
||||
"""
|
||||
success = ProjectManager.delete_project(project_id)
|
||||
project = _scoped_project(project_id)
|
||||
if project is None:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.projectDeleteFailed', id=project_id)
|
||||
}), 404
|
||||
|
||||
actor = current_actor()
|
||||
success = ProjectManager.delete_project(
|
||||
project.project_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id if actor.role is Role.USER else None,
|
||||
)
|
||||
|
||||
if not success:
|
||||
return jsonify({
|
||||
@@ -87,11 +188,12 @@ def delete_project(project_id: str):
|
||||
|
||||
|
||||
@graph_bp.route('/project/<project_id>/reset', methods=['POST'])
|
||||
@require_auth
|
||||
def reset_project(project_id: str):
|
||||
"""
|
||||
重置项目状态(用于重新构建图谱)
|
||||
"""
|
||||
project = ProjectManager.get_project(project_id)
|
||||
project = _scoped_project(project_id)
|
||||
|
||||
if not project:
|
||||
return jsonify({
|
||||
@@ -120,6 +222,8 @@ def reset_project(project_id: str):
|
||||
# ============== 接口1:上传文件并生成本体 ==============
|
||||
|
||||
@graph_bp.route('/ontology/generate', methods=['POST'])
|
||||
@require_auth
|
||||
@idempotent
|
||||
def generate_ontology():
|
||||
"""
|
||||
接口1:上传文件,分析生成本体定义
|
||||
@@ -190,9 +294,16 @@ def generate_ontology():
|
||||
"error": t('api.requireFileUpload')
|
||||
}), 400
|
||||
|
||||
# 创建项目
|
||||
project = ProjectManager.create_project(name=project_name)
|
||||
# Create a tenant-owned project.
|
||||
actor = current_actor()
|
||||
project = ProjectManager.create_project(
|
||||
name=project_name,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id,
|
||||
)
|
||||
project.simulation_requirement = simulation_requirement
|
||||
# Best-effort dual-write into the durable repository.
|
||||
_sync_project_to_durable(project)
|
||||
logger.info(f"创建项目: {project.project_id}")
|
||||
|
||||
# 保存文件并提取文本
|
||||
@@ -219,7 +330,11 @@ def generate_ontology():
|
||||
all_text += f"\n\n=== {file_info['original_filename']} ===\n{text}"
|
||||
|
||||
if not document_texts:
|
||||
ProjectManager.delete_project(project.project_id)
|
||||
ProjectManager.delete_project(
|
||||
project.project_id,
|
||||
organization_id=current_actor().organization_id,
|
||||
owner_user_id=current_actor().user_id,
|
||||
)
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.noDocProcessed')
|
||||
@@ -252,6 +367,7 @@ def generate_ontology():
|
||||
project.analysis_summary = ontology.get("analysis_summary", "")
|
||||
project.status = ProjectStatus.ONTOLOGY_GENERATED
|
||||
ProjectManager.save_project(project)
|
||||
_sync_project_to_durable(project)
|
||||
logger.info(f"=== 本体生成完成 === 项目ID: {project.project_id}")
|
||||
|
||||
return jsonify({
|
||||
@@ -267,16 +383,15 @@ def generate_ontology():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
logger.exception("Graph operation failed: %s", type(e).__name__)
|
||||
raise ApiError("graph_operation_failed", 500, "api.internalError") from e
|
||||
|
||||
|
||||
# ============== 接口2:构建图谱 ==============
|
||||
|
||||
@graph_bp.route('/build', methods=['POST'])
|
||||
@require_auth
|
||||
@idempotent
|
||||
def build_graph():
|
||||
"""
|
||||
接口2:根据project_id构建图谱
|
||||
@@ -302,12 +417,17 @@ def build_graph():
|
||||
try:
|
||||
logger.info("=== 开始构建图谱 ===")
|
||||
|
||||
# 检查配置
|
||||
# Validate the selected backend before starting an asynchronous task.
|
||||
backend = Config.MEMORY_BACKEND
|
||||
if backend not in {"zep", "local"}:
|
||||
raise ApiError("invalid_memory_backend", 500, "api.internalError")
|
||||
errors = []
|
||||
if not Config.ZEP_API_KEY:
|
||||
if backend == "zep" and not Config.ZEP_API_KEY:
|
||||
errors.append(t('api.zepApiKeyMissing'))
|
||||
if backend == "local" and not Config.LLM_API_KEY:
|
||||
errors.append("LLM_API_KEY not configured")
|
||||
if errors:
|
||||
logger.error(f"配置错误: {errors}")
|
||||
logger.error("Graph backend configuration is incomplete: backend=%s", backend)
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.configError', details="; ".join(errors))
|
||||
@@ -325,7 +445,7 @@ def build_graph():
|
||||
}), 400
|
||||
|
||||
# 获取项目
|
||||
project = ProjectManager.get_project(project_id)
|
||||
project = _scoped_project(project_id)
|
||||
if not project:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
@@ -380,15 +500,28 @@ def build_graph():
|
||||
"error": t('api.ontologyNotFound')
|
||||
}), 400
|
||||
|
||||
# 创建异步任务
|
||||
# Create durable task metadata with tenant scope.
|
||||
task_manager = TaskManager()
|
||||
task_id = task_manager.create_task(f"构建图谱: {graph_name}")
|
||||
actor = current_actor()
|
||||
organization_id = actor.organization_id
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
raise ApiError("memory_backend_unavailable", 503, "api.internalError")
|
||||
task_id = task_manager.create_task(
|
||||
f"构建图谱: {graph_name}",
|
||||
metadata={
|
||||
"project_id": project_id,
|
||||
"organization_id": actor.organization_id,
|
||||
"owner_user_id": actor.user_id,
|
||||
},
|
||||
)
|
||||
logger.info(f"创建图谱构建任务: task_id={task_id}, project_id={project_id}")
|
||||
|
||||
# 更新项目状态
|
||||
project.status = ProjectStatus.GRAPH_BUILDING
|
||||
project.graph_build_task_id = task_id
|
||||
ProjectManager.save_project(project)
|
||||
_sync_project_to_durable(project)
|
||||
|
||||
# Capture locale before spawning background thread
|
||||
current_locale = get_locale()
|
||||
@@ -405,8 +538,22 @@ def build_graph():
|
||||
message=t('progress.initGraphService')
|
||||
)
|
||||
|
||||
# 创建图谱构建服务
|
||||
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
|
||||
# Choose the storage adapter inside the worker with captured scope.
|
||||
if backend == "local":
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id=organization_id,
|
||||
project_id=project_id,
|
||||
extraction_client=LLMClient(),
|
||||
language=current_locale,
|
||||
)
|
||||
else:
|
||||
builder = GraphBuilderService(
|
||||
api_key=Config.ZEP_API_KEY,
|
||||
organization_id=organization_id,
|
||||
owner_user_id=actor.user_id,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
|
||||
# 分块
|
||||
task_manager.update_task(
|
||||
@@ -424,7 +571,7 @@ def build_graph():
|
||||
# 创建图谱
|
||||
task_manager.update_task(
|
||||
task_id,
|
||||
message=t('progress.creatingZepGraph'),
|
||||
message=t('progress.creatingGraph' if backend == "local" else 'progress.creatingZepGraph'),
|
||||
progress=10
|
||||
)
|
||||
graph_id = builder.create_graph(name=graph_name)
|
||||
@@ -432,6 +579,7 @@ def build_graph():
|
||||
# 更新项目的graph_id
|
||||
project.graph_id = graph_id
|
||||
ProjectManager.save_project(project)
|
||||
_sync_project_to_durable(project, session_factory)
|
||||
|
||||
# 设置本体
|
||||
task_manager.update_task(
|
||||
@@ -466,7 +614,7 @@ def build_graph():
|
||||
# 等待Zep处理完成(查询每个episode的processed状态)
|
||||
task_manager.update_task(
|
||||
task_id,
|
||||
message=t('progress.waitingZepProcess'),
|
||||
message=t('progress.processingComplete' if backend == "local" else 'progress.waitingZepProcess'),
|
||||
progress=55
|
||||
)
|
||||
|
||||
@@ -491,6 +639,7 @@ def build_graph():
|
||||
# 更新项目状态
|
||||
project.status = ProjectStatus.GRAPH_COMPLETED
|
||||
ProjectManager.save_project(project)
|
||||
_sync_project_to_durable(project, session_factory)
|
||||
|
||||
node_count = graph_data.get("node_count", 0)
|
||||
edge_count = graph_data.get("edge_count", 0)
|
||||
@@ -513,20 +662,24 @@ def build_graph():
|
||||
|
||||
except Exception as e:
|
||||
# 更新项目状态为失败
|
||||
build_logger.error(f"[{task_id}] 图谱构建失败: {str(e)}")
|
||||
build_logger.debug(traceback.format_exc())
|
||||
build_logger.error(
|
||||
"[%s] graph build failed: error_type=%s",
|
||||
task_id,
|
||||
type(e).__name__,
|
||||
)
|
||||
|
||||
project.status = ProjectStatus.FAILED
|
||||
project.error = str(e)
|
||||
project.error = t('api.internalError')
|
||||
ProjectManager.save_project(project)
|
||||
_sync_project_to_durable(project, session_factory)
|
||||
|
||||
task_manager.update_task(
|
||||
task_id,
|
||||
status=TaskStatus.FAILED,
|
||||
message=t('progress.buildFailed', error=str(e)),
|
||||
error=traceback.format_exc()
|
||||
message=t('api.internalError'),
|
||||
error=t('api.internalError'),
|
||||
)
|
||||
|
||||
|
||||
# 启动后台线程
|
||||
thread = threading.Thread(target=build_task, daemon=True)
|
||||
thread.start()
|
||||
@@ -541,21 +694,19 @@ def build_graph():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
logger.exception("Graph operation failed: %s", type(e).__name__)
|
||||
raise ApiError("graph_operation_failed", 500, "api.internalError") from e
|
||||
|
||||
|
||||
# ============== 任务查询接口 ==============
|
||||
|
||||
@graph_bp.route('/task/<task_id>', methods=['GET'])
|
||||
@require_auth
|
||||
def get_task(task_id: str):
|
||||
"""
|
||||
查询任务状态
|
||||
"""
|
||||
task = TaskManager().get_task(task_id)
|
||||
task = _scoped_task(task_id)
|
||||
|
||||
if not task:
|
||||
return jsonify({
|
||||
@@ -570,11 +721,19 @@ def get_task(task_id: str):
|
||||
|
||||
|
||||
@graph_bp.route('/tasks', methods=['GET'])
|
||||
@require_auth
|
||||
def list_tasks():
|
||||
"""
|
||||
列出所有任务
|
||||
"""
|
||||
tasks = TaskManager().list_tasks()
|
||||
actor = current_actor()
|
||||
tasks = [
|
||||
task for task in TaskManager().list_tasks(
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id if actor.role is Role.USER else None,
|
||||
)
|
||||
if _scoped_task(task.task_id) is not None
|
||||
]
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -586,18 +745,29 @@ def list_tasks():
|
||||
# ============== 图谱数据接口 ==============
|
||||
|
||||
@graph_bp.route('/data/<graph_id>', methods=['GET'])
|
||||
@require_auth
|
||||
def get_graph_data(graph_id: str):
|
||||
"""
|
||||
获取图谱数据(节点和边)
|
||||
"""
|
||||
project = _scoped_graph(graph_id)
|
||||
if project is None:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.graphNotBuilt'),
|
||||
}), 404
|
||||
try:
|
||||
if not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
}), 500
|
||||
|
||||
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
builder = _local_graph_builder(project)
|
||||
elif Config.MEMORY_BACKEND == "zep":
|
||||
if not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
}), 500
|
||||
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
|
||||
else:
|
||||
raise ApiError("invalid_memory_backend", 500, "api.internalError")
|
||||
graph_data = builder.get_graph_data(graph_id)
|
||||
|
||||
return jsonify({
|
||||
@@ -606,26 +776,34 @@ def get_graph_data(graph_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
logger.exception("Graph operation failed: %s", type(e).__name__)
|
||||
raise ApiError("graph_operation_failed", 500, "api.internalError") from e
|
||||
|
||||
|
||||
@graph_bp.route('/delete/<graph_id>', methods=['DELETE'])
|
||||
@require_auth
|
||||
def delete_graph(graph_id: str):
|
||||
"""
|
||||
删除Zep图谱
|
||||
ลบกราฟ memory
|
||||
"""
|
||||
project = _scoped_graph(graph_id)
|
||||
if project is None:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.graphNotBuilt'),
|
||||
}), 404
|
||||
try:
|
||||
if not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
}), 500
|
||||
|
||||
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
builder = _local_graph_builder(project)
|
||||
elif Config.MEMORY_BACKEND == "zep":
|
||||
if not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
}), 500
|
||||
builder = GraphBuilderService(api_key=Config.ZEP_API_KEY)
|
||||
else:
|
||||
raise ApiError("invalid_memory_backend", 500, "api.internalError")
|
||||
builder.delete_graph(graph_id)
|
||||
|
||||
return jsonify({
|
||||
@@ -634,8 +812,5 @@ def delete_graph(graph_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
logger.exception("Graph operation failed: %s", type(e).__name__)
|
||||
raise ApiError("graph_operation_failed", 500, "api.internalError") from e
|
||||
|
||||
@@ -4,25 +4,107 @@ Report API路由
|
||||
"""
|
||||
|
||||
import os
|
||||
import traceback
|
||||
import threading
|
||||
from flask import request, jsonify, send_file
|
||||
from flask import current_app, request, jsonify, send_file
|
||||
|
||||
from . import report_bp
|
||||
from ..config import Config
|
||||
from ..services.report_agent import ReportAgent, ReportManager, ReportStatus
|
||||
from ..services.simulation_manager import SimulationManager
|
||||
from ..models.project import ProjectManager
|
||||
from ..models.task import TaskManager, TaskStatus
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import t, get_locale, set_locale
|
||||
from ..security.auth import authenticate_readonly_request, current_actor
|
||||
from ..security.resources import enforce_request_scope, scoped_reports, require_scoped_project, require_scoped_simulation
|
||||
from ..services.memory_tools import LocalMemoryTools
|
||||
from ..services.product_repository import ProductRepository
|
||||
from ..services.idempotency import idempotent
|
||||
from ..utils.api_errors import ApiError, internal_error_payload
|
||||
|
||||
logger = get_logger('crowdsight.api.report')
|
||||
|
||||
|
||||
@report_bp.errorhandler(ApiError)
|
||||
def handle_report_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
|
||||
def _safe_internal_failure(operation: str, error: Exception):
|
||||
"""Return a safe response without leaking exception details."""
|
||||
if isinstance(error, ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
logger.error("%s failed: error_type=%s", operation, type(error).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
|
||||
def _sync_report_to_durable(
|
||||
report,
|
||||
*,
|
||||
organization_id,
|
||||
project_id,
|
||||
simulation_id,
|
||||
created_by_user_id,
|
||||
session_factory=None,
|
||||
):
|
||||
"""Best-effort dual-write of a legacy report into durable SQL.
|
||||
|
||||
Never raises; the filesystem ReportManager remains authoritative if the
|
||||
durable write is unavailable. Callers inside a worker thread pass the
|
||||
captured ``session_factory`` explicitly.
|
||||
"""
|
||||
try:
|
||||
if session_factory is None:
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
return
|
||||
session = session_factory()
|
||||
try:
|
||||
ProductRepository(session).sync_report(
|
||||
report,
|
||||
organization_id=organization_id,
|
||||
project_id=project_id,
|
||||
simulation_id=simulation_id,
|
||||
created_by_user_id=created_by_user_id,
|
||||
commit=True,
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
except Exception:
|
||||
logger.warning("durable report sync skipped", exc_info=True)
|
||||
|
||||
|
||||
@report_bp.before_request
|
||||
def _authenticate_report_request():
|
||||
authenticate_readonly_request()
|
||||
enforce_request_scope()
|
||||
|
||||
|
||||
def _request_memory_tools(graph_id: str):
|
||||
"""Build a request-scoped local memory adapter when local mode is active."""
|
||||
if Config.MEMORY_BACKEND != "local":
|
||||
return None, None
|
||||
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
raise ApiError("memory_backend_unavailable", 503, "api.internalError")
|
||||
|
||||
actor = current_actor()
|
||||
session = session_factory()
|
||||
try:
|
||||
tools = LocalMemoryTools(
|
||||
session,
|
||||
organization_id=actor.organization_id,
|
||||
graph_id=graph_id,
|
||||
)
|
||||
except Exception:
|
||||
session.close()
|
||||
raise
|
||||
return tools, session
|
||||
|
||||
|
||||
# ============== 报告生成接口 ==============
|
||||
|
||||
@report_bp.route('/generate', methods=['POST'])
|
||||
@idempotent
|
||||
def generate_report():
|
||||
"""
|
||||
生成模拟分析报告(异步任务)
|
||||
@@ -59,15 +141,8 @@ def generate_report():
|
||||
|
||||
force_regenerate = data.get('force_regenerate', False)
|
||||
|
||||
# 获取模拟信息
|
||||
manager = SimulationManager()
|
||||
state = manager.get_simulation(simulation_id)
|
||||
|
||||
if not state:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.simulationNotFound', id=simulation_id)
|
||||
}), 404
|
||||
# Fetch both resources through the tenant/owner scope guards.
|
||||
state = require_scoped_simulation(simulation_id)
|
||||
|
||||
# 检查是否已有报告
|
||||
if not force_regenerate:
|
||||
@@ -85,12 +160,7 @@ def generate_report():
|
||||
})
|
||||
|
||||
# 获取项目信息
|
||||
project = ProjectManager.get_project(state.project_id)
|
||||
if not project:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.projectNotFound', id=state.project_id)
|
||||
}), 404
|
||||
project = require_scoped_project(state.project_id)
|
||||
|
||||
graph_id = state.graph_id or project.graph_id
|
||||
if not graph_id:
|
||||
@@ -110,14 +180,26 @@ def generate_report():
|
||||
import uuid
|
||||
report_id = f"report_{uuid.uuid4().hex[:12]}"
|
||||
|
||||
# 创建异步任务
|
||||
# Create an ownership-scoped in-memory task record.
|
||||
task_manager = TaskManager()
|
||||
actor = current_actor()
|
||||
organization_id = actor.organization_id
|
||||
owner_user_id = actor.user_id
|
||||
project_id = state.project_id
|
||||
memory_backend = Config.MEMORY_BACKEND
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory") if memory_backend == "local" else None
|
||||
if memory_backend not in {"zep", "local"}:
|
||||
raise ApiError("invalid_memory_backend", 500, "api.internalError")
|
||||
if memory_backend == "local" and session_factory is None:
|
||||
raise ApiError("memory_backend_unavailable", 503, "api.internalError")
|
||||
task_id = task_manager.create_task(
|
||||
task_type="report_generate",
|
||||
metadata={
|
||||
"simulation_id": simulation_id,
|
||||
"graph_id": graph_id,
|
||||
"report_id": report_id
|
||||
"report_id": report_id,
|
||||
"organization_id": actor.organization_id,
|
||||
"owner_user_id": actor.user_id,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -127,6 +209,7 @@ def generate_report():
|
||||
# 定义后台任务
|
||||
def run_generate():
|
||||
set_locale(current_locale)
|
||||
local_session = None
|
||||
try:
|
||||
task_manager.update_task(
|
||||
task_id,
|
||||
@@ -134,12 +217,23 @@ def generate_report():
|
||||
progress=0,
|
||||
message=t('api.initReportAgent')
|
||||
)
|
||||
|
||||
# 创建Report Agent
|
||||
|
||||
memory_tools = None
|
||||
if memory_backend == "local":
|
||||
if session_factory is None:
|
||||
raise RuntimeError("memory_backend_unavailable")
|
||||
local_session = session_factory()
|
||||
memory_tools = LocalMemoryTools(
|
||||
local_session,
|
||||
organization_id=organization_id,
|
||||
graph_id=graph_id,
|
||||
)
|
||||
|
||||
agent = ReportAgent(
|
||||
graph_id=graph_id,
|
||||
simulation_id=simulation_id,
|
||||
simulation_requirement=simulation_requirement
|
||||
simulation_requirement=simulation_requirement,
|
||||
memory_tools=memory_tools,
|
||||
)
|
||||
|
||||
# 进度回调
|
||||
@@ -158,6 +252,15 @@ def generate_report():
|
||||
|
||||
# 保存报告
|
||||
ReportManager.save_report(report)
|
||||
# Best-effort dual-write into the durable repository.
|
||||
_sync_report_to_durable(
|
||||
report,
|
||||
organization_id=organization_id,
|
||||
project_id=project_id,
|
||||
simulation_id=simulation_id,
|
||||
created_by_user_id=owner_user_id,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
|
||||
if report.status == ReportStatus.COMPLETED:
|
||||
task_manager.complete_task(
|
||||
@@ -169,11 +272,14 @@ def generate_report():
|
||||
}
|
||||
)
|
||||
else:
|
||||
task_manager.fail_task(task_id, report.error or t('api.reportGenerateFailed'))
|
||||
task_manager.fail_task(task_id, t('api.reportGenerateFailed'))
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"报告生成失败: {str(e)}")
|
||||
task_manager.fail_task(task_id, str(e))
|
||||
except Exception as exc:
|
||||
logger.error("Report generation failed: error=%s", type(exc).__name__)
|
||||
task_manager.fail_task(task_id, t('api.reportGenerateFailed'))
|
||||
finally:
|
||||
if local_session is not None:
|
||||
local_session.close()
|
||||
|
||||
# 启动后台线程
|
||||
thread = threading.Thread(target=run_generate, daemon=True)
|
||||
@@ -192,12 +298,7 @@ def generate_report():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"启动报告生成任务失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("start report generation", e)
|
||||
|
||||
|
||||
@report_bp.route('/generate/status', methods=['POST'])
|
||||
@@ -265,11 +366,7 @@ def get_generate_status():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"查询任务状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 500
|
||||
return _safe_internal_failure("get report task status", e)
|
||||
|
||||
|
||||
# ============== 报告获取接口 ==============
|
||||
@@ -308,12 +405,7 @@ def get_report(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取报告失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get report", e)
|
||||
|
||||
|
||||
@report_bp.route('/by-simulation/<simulation_id>', methods=['GET'])
|
||||
@@ -347,12 +439,7 @@ def get_report_by_simulation(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取报告失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get report", e)
|
||||
|
||||
|
||||
@report_bp.route('/list', methods=['GET'])
|
||||
@@ -375,9 +462,9 @@ def list_reports():
|
||||
simulation_id = request.args.get('simulation_id')
|
||||
limit = request.args.get('limit', 50, type=int)
|
||||
|
||||
reports = ReportManager.list_reports(
|
||||
reports = scoped_reports(
|
||||
simulation_id=simulation_id,
|
||||
limit=limit
|
||||
limit=limit,
|
||||
)
|
||||
|
||||
return jsonify({
|
||||
@@ -387,12 +474,7 @@ def list_reports():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"列出报告失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("list reports", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>/download', methods=['GET'])
|
||||
@@ -433,12 +515,7 @@ def download_report(report_id: str):
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"下载报告失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("download report", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>', methods=['DELETE'])
|
||||
@@ -459,17 +536,13 @@ def delete_report(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"删除报告失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("delete report", e)
|
||||
|
||||
|
||||
# ============== Report Agent对话接口 ==============
|
||||
|
||||
@report_bp.route('/chat', methods=['POST'])
|
||||
@idempotent
|
||||
def chat_with_report_agent():
|
||||
"""
|
||||
与Report Agent对话
|
||||
@@ -515,22 +588,9 @@ def chat_with_report_agent():
|
||||
"error": t('api.requireMessage')
|
||||
}), 400
|
||||
|
||||
# 获取模拟和项目信息
|
||||
manager = SimulationManager()
|
||||
state = manager.get_simulation(simulation_id)
|
||||
|
||||
if not state:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.simulationNotFound', id=simulation_id)
|
||||
}), 404
|
||||
|
||||
project = ProjectManager.get_project(state.project_id)
|
||||
if not project:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.projectNotFound', id=state.project_id)
|
||||
}), 404
|
||||
# Resolve both resources through tenant/owner scope guards.
|
||||
state = require_scoped_simulation(simulation_id)
|
||||
project = require_scoped_project(state.project_id)
|
||||
|
||||
graph_id = state.graph_id or project.graph_id
|
||||
if not graph_id:
|
||||
@@ -541,27 +601,25 @@ def chat_with_report_agent():
|
||||
|
||||
simulation_requirement = project.simulation_requirement or ""
|
||||
|
||||
# 创建Agent并进行对话
|
||||
agent = ReportAgent(
|
||||
graph_id=graph_id,
|
||||
simulation_id=simulation_id,
|
||||
simulation_requirement=simulation_requirement
|
||||
)
|
||||
|
||||
result = agent.chat(message=message, chat_history=chat_history)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"data": result
|
||||
})
|
||||
memory_tools, local_session = _request_memory_tools(graph_id)
|
||||
try:
|
||||
agent = ReportAgent(
|
||||
graph_id=graph_id,
|
||||
simulation_id=simulation_id,
|
||||
simulation_requirement=simulation_requirement,
|
||||
memory_tools=memory_tools,
|
||||
)
|
||||
result = agent.chat(message=message, chat_history=chat_history)
|
||||
return jsonify({
|
||||
"success": True,
|
||||
"data": result
|
||||
})
|
||||
finally:
|
||||
if local_session is not None:
|
||||
local_session.close()
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"对话失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("chat with report agent", e)
|
||||
|
||||
|
||||
# ============== 报告进度与分章节接口 ==============
|
||||
@@ -599,12 +657,7 @@ def get_report_progress(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取报告进度失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get report progress", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>/sections', methods=['GET'])
|
||||
@@ -650,12 +703,7 @@ def get_report_sections(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取章节列表失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get report sections", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>/section/<int:section_index>', methods=['GET'])
|
||||
@@ -694,12 +742,7 @@ def get_single_section(report_id: str, section_index: int):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取章节内容失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get report section", e)
|
||||
|
||||
|
||||
# ============== 报告状态检查接口 ==============
|
||||
@@ -745,12 +788,7 @@ def check_report_status(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"检查报告状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("check report status", e)
|
||||
|
||||
|
||||
# ============== Agent 日志接口 ==============
|
||||
@@ -806,12 +844,7 @@ def get_agent_log(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取Agent日志失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get agent log", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>/agent-log/stream', methods=['GET'])
|
||||
@@ -840,12 +873,7 @@ def stream_agent_log(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取Agent日志失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get agent log", e)
|
||||
|
||||
|
||||
# ============== 控制台日志接口 ==============
|
||||
@@ -888,12 +916,7 @@ def get_console_log(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取控制台日志失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get console log", e)
|
||||
|
||||
|
||||
@report_bp.route('/<report_id>/console-log/stream', methods=['GET'])
|
||||
@@ -922,12 +945,7 @@ def stream_console_log(report_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取控制台日志失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get console log", e)
|
||||
|
||||
|
||||
# ============== 工具调用接口(供调试使用)==============
|
||||
@@ -957,14 +975,19 @@ def search_graph_tool():
|
||||
"error": t('api.requireGraphIdAndQuery')
|
||||
}), 400
|
||||
|
||||
from ..services.zep_tools import ZepToolsService
|
||||
|
||||
tools = ZepToolsService()
|
||||
result = tools.search_graph(
|
||||
graph_id=graph_id,
|
||||
query=query,
|
||||
limit=limit
|
||||
)
|
||||
tools, local_session = _request_memory_tools(graph_id)
|
||||
if tools is None:
|
||||
from ..services.zep_tools import ZepToolsService
|
||||
tools = ZepToolsService()
|
||||
try:
|
||||
result = tools.search_graph(
|
||||
graph_id=graph_id,
|
||||
query=query,
|
||||
limit=limit
|
||||
)
|
||||
finally:
|
||||
if local_session is not None:
|
||||
local_session.close()
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -972,12 +995,7 @@ def search_graph_tool():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"图谱搜索失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("search graph tool", e)
|
||||
|
||||
|
||||
@report_bp.route('/tools/statistics', methods=['POST'])
|
||||
@@ -1001,10 +1019,15 @@ def get_graph_statistics_tool():
|
||||
"error": t('api.requireGraphId')
|
||||
}), 400
|
||||
|
||||
from ..services.zep_tools import ZepToolsService
|
||||
|
||||
tools = ZepToolsService()
|
||||
result = tools.get_graph_statistics(graph_id)
|
||||
tools, local_session = _request_memory_tools(graph_id)
|
||||
if tools is None:
|
||||
from ..services.zep_tools import ZepToolsService
|
||||
tools = ZepToolsService()
|
||||
try:
|
||||
result = tools.get_graph_statistics(graph_id)
|
||||
finally:
|
||||
if local_session is not None:
|
||||
local_session.close()
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -1012,9 +1035,4 @@ def get_graph_statistics_tool():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取图谱统计失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("get graph statistics tool", e)
|
||||
|
||||
@@ -5,22 +5,100 @@ Step2: Zep实体读取与过滤、OASIS模拟准备与运行(全程自动化
|
||||
|
||||
import os
|
||||
import json
|
||||
import traceback
|
||||
from flask import request, jsonify, send_file
|
||||
from flask import current_app, request, jsonify, send_file
|
||||
|
||||
from . import simulation_bp
|
||||
from ..config import Config
|
||||
from ..services.zep_entity_reader import ZepEntityReader
|
||||
from ..services.oasis_profile_generator import OasisProfileGenerator
|
||||
from ..services.simulation_manager import SimulationManager, SimulationStatus
|
||||
from ..services.simulation_runner import SimulationRunner, RunnerStatus
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import t, get_locale, set_locale
|
||||
from ..models.project import ProjectManager
|
||||
from ..security.auth import authenticate_readonly_request, current_actor
|
||||
from ..security.resources import (
|
||||
enforce_request_scope,
|
||||
require_scoped_simulation,
|
||||
scoped_project,
|
||||
scoped_simulations,
|
||||
)
|
||||
from ..services.memory_entity_reader import make_local_entity_reader_factory
|
||||
from ..services.memory_tools import LocalMemoryTools
|
||||
from ..services.product_repository import ProductRepository
|
||||
from ..services.idempotency import idempotent
|
||||
from ..utils.api_errors import ApiError, internal_error_payload
|
||||
|
||||
logger = get_logger('crowdsight.api.simulation')
|
||||
|
||||
|
||||
def _safe_internal_failure(operation: str, error: Exception):
|
||||
if isinstance(error, ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
logger.error("%s failed: error_type=%s", operation, type(error).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
|
||||
def _safe_client_failure(status_code: int, code: str = "invalid_request"):
|
||||
return jsonify(
|
||||
ApiError(code, status_code, "api.requestError").to_payload(t)
|
||||
), status_code
|
||||
|
||||
|
||||
def _sync_simulation_to_durable(state, *, organization_id, project_id, created_by_user_id):
|
||||
"""Best-effort dual-write of a legacy simulation state into durable SQL.
|
||||
|
||||
Never raises; the filesystem manager remains authoritative if the durable
|
||||
write is unavailable. The SimulationState carries no tenant id, so the
|
||||
caller supplies organization/project/creator which are merged in.
|
||||
"""
|
||||
try:
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
return
|
||||
payload = state.to_dict() if hasattr(state, "to_dict") else dict(state)
|
||||
payload["organization_id"] = organization_id
|
||||
payload["project_id"] = project_id
|
||||
payload["created_by_user_id"] = created_by_user_id
|
||||
session = session_factory()
|
||||
try:
|
||||
ProductRepository(session).sync_simulation(payload, commit=True)
|
||||
finally:
|
||||
session.close()
|
||||
except Exception:
|
||||
logger.warning("durable simulation sync skipped", exc_info=True)
|
||||
|
||||
|
||||
def _simulation_manager_for_request() -> SimulationManager:
|
||||
backend = Config.MEMORY_BACKEND
|
||||
if backend == "zep":
|
||||
return SimulationManager()
|
||||
if backend != "local":
|
||||
raise ApiError("invalid_memory_backend", 500, "api.internalError")
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
raise ApiError("memory_backend_unavailable", 503, "api.internalError")
|
||||
reader_factory = make_local_entity_reader_factory(
|
||||
session_factory,
|
||||
organization_id=current_actor().organization_id,
|
||||
)
|
||||
return SimulationManager(entity_reader_factory=reader_factory)
|
||||
|
||||
|
||||
def _entity_reader_for_request(graph_id: str):
|
||||
return _simulation_manager_for_request().create_entity_reader(graph_id)
|
||||
|
||||
|
||||
@simulation_bp.errorhandler(ApiError)
|
||||
def handle_simulation_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
|
||||
@simulation_bp.before_request
|
||||
def _authenticate_simulation_request():
|
||||
authenticate_readonly_request()
|
||||
enforce_request_scope()
|
||||
|
||||
|
||||
# Interview prompt 优化前缀
|
||||
# 添加此前缀可以避免Agent调用工具,直接用文本回复
|
||||
INTERVIEW_PROMPT_PREFIX = "结合你的人设、所有的过往记忆与行动,不调用任何工具直接用文本回复我:"
|
||||
@@ -58,7 +136,7 @@ def get_graph_entities(graph_id: str):
|
||||
enrich: 是否获取相关边信息(默认true)
|
||||
"""
|
||||
try:
|
||||
if not Config.ZEP_API_KEY:
|
||||
if Config.MEMORY_BACKEND == "zep" and not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
@@ -70,12 +148,17 @@ def get_graph_entities(graph_id: str):
|
||||
|
||||
logger.info(f"获取图谱实体: graph_id={graph_id}, entity_types={entity_types}, enrich={enrich}")
|
||||
|
||||
reader = ZepEntityReader()
|
||||
result = reader.filter_defined_entities(
|
||||
graph_id=graph_id,
|
||||
defined_entity_types=entity_types,
|
||||
enrich_with_edges=enrich
|
||||
)
|
||||
reader = _entity_reader_for_request(graph_id)
|
||||
try:
|
||||
result = reader.filter_defined_entities(
|
||||
graph_id=graph_id,
|
||||
defined_entity_types=entity_types,
|
||||
enrich_with_edges=enrich,
|
||||
)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -83,26 +166,26 @@ def get_graph_entities(graph_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取图谱实体失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/entities/<graph_id>/<entity_uuid>', methods=['GET'])
|
||||
def get_entity_detail(graph_id: str, entity_uuid: str):
|
||||
"""获取单个实体的详细信息"""
|
||||
try:
|
||||
if not Config.ZEP_API_KEY:
|
||||
if Config.MEMORY_BACKEND == "zep" and not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
}), 500
|
||||
|
||||
reader = ZepEntityReader()
|
||||
entity = reader.get_entity_with_context(graph_id, entity_uuid)
|
||||
reader = _entity_reader_for_request(graph_id)
|
||||
try:
|
||||
entity = reader.get_entity_with_context(graph_id, entity_uuid)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
|
||||
if not entity:
|
||||
return jsonify({
|
||||
@@ -116,19 +199,14 @@ def get_entity_detail(graph_id: str, entity_uuid: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取实体详情失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/entities/<graph_id>/by-type/<entity_type>', methods=['GET'])
|
||||
def get_entities_by_type(graph_id: str, entity_type: str):
|
||||
"""获取指定类型的所有实体"""
|
||||
try:
|
||||
if not Config.ZEP_API_KEY:
|
||||
if Config.MEMORY_BACKEND == "zep" and not Config.ZEP_API_KEY:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.zepApiKeyMissing')
|
||||
@@ -136,12 +214,17 @@ def get_entities_by_type(graph_id: str, entity_type: str):
|
||||
|
||||
enrich = request.args.get('enrich', 'true').lower() == 'true'
|
||||
|
||||
reader = ZepEntityReader()
|
||||
entities = reader.get_entities_by_type(
|
||||
graph_id=graph_id,
|
||||
entity_type=entity_type,
|
||||
enrich_with_edges=enrich
|
||||
)
|
||||
reader = _entity_reader_for_request(graph_id)
|
||||
try:
|
||||
entities = reader.get_entities_by_type(
|
||||
graph_id=graph_id,
|
||||
entity_type=entity_type,
|
||||
enrich_with_edges=enrich,
|
||||
)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -153,17 +236,13 @@ def get_entities_by_type(graph_id: str, entity_type: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取实体失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== 模拟管理接口 ==============
|
||||
|
||||
@simulation_bp.route('/create', methods=['POST'])
|
||||
@idempotent
|
||||
def create_simulation():
|
||||
"""
|
||||
创建新的模拟
|
||||
@@ -202,27 +281,51 @@ def create_simulation():
|
||||
"error": t('api.requireProjectId')
|
||||
}), 400
|
||||
|
||||
project = ProjectManager.get_project(project_id)
|
||||
project = scoped_project(project_id)
|
||||
if not project:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.projectNotFound', id=project_id)
|
||||
}), 404
|
||||
|
||||
|
||||
graph_id = data.get('graph_id') or project.graph_id
|
||||
if not graph_id:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.graphNotBuilt')
|
||||
}), 400
|
||||
|
||||
actor = current_actor()
|
||||
graph_project = ProjectManager.find_project_by_graph_id(
|
||||
graph_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id,
|
||||
)
|
||||
if graph_project is None or graph_project.project_id != project.project_id:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.graphNotBuilt')
|
||||
}), 404
|
||||
|
||||
manager = SimulationManager()
|
||||
manager = SimulationManager(
|
||||
session_factory=current_app.extensions.get("crowdsight_session_factory")
|
||||
)
|
||||
state = manager.create_simulation(
|
||||
project_id=project_id,
|
||||
graph_id=graph_id,
|
||||
enable_twitter=data.get('enable_twitter', True),
|
||||
enable_reddit=data.get('enable_reddit', True),
|
||||
)
|
||||
# Attach tenant scope so every subsequently saved state mirrors durable.
|
||||
state.organization_id = actor.organization_id
|
||||
state.owner_user_id = actor.user_id
|
||||
# Best-effort dual-write into the durable repository.
|
||||
_sync_simulation_to_durable(
|
||||
state,
|
||||
organization_id=actor.organization_id,
|
||||
project_id=project_id,
|
||||
created_by_user_id=actor.user_id,
|
||||
)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -230,12 +333,7 @@ def create_simulation():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"创建模拟失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
def _check_simulation_prepared(simulation_id: str) -> tuple:
|
||||
@@ -332,7 +430,7 @@ def _check_simulation_prepared(simulation_id: str) -> tuple:
|
||||
logger.info(f"自动更新模拟状态: {simulation_id} preparing -> ready")
|
||||
status = "ready"
|
||||
except Exception as e:
|
||||
logger.warning(f"自动更新状态失败: {e}")
|
||||
logger.warning("自动更新状态失败: error_type=%s", type(e).__name__)
|
||||
|
||||
logger.info(f"模拟 {simulation_id} 检测结果: 已准备完成 (status={status}, config_generated={config_generated})")
|
||||
return True, {
|
||||
@@ -354,10 +452,12 @@ def _check_simulation_prepared(simulation_id: str) -> tuple:
|
||||
}
|
||||
|
||||
except Exception as e:
|
||||
return False, {"reason": f"读取状态文件失败: {str(e)}"}
|
||||
logger.warning("读取状态文件失败: error_type=%s", type(e).__name__)
|
||||
return False, {"reason": "preparation_status_unavailable"}
|
||||
|
||||
|
||||
@simulation_bp.route('/prepare', methods=['POST'])
|
||||
@idempotent
|
||||
def prepare_simulation():
|
||||
"""
|
||||
准备模拟环境(异步任务,LLM智能生成所有参数)
|
||||
@@ -413,7 +513,7 @@ def prepare_simulation():
|
||||
"error": t('api.requireSimulationId')
|
||||
}), 400
|
||||
|
||||
manager = SimulationManager()
|
||||
manager = _simulation_manager_for_request()
|
||||
state = manager.get_simulation(simulation_id)
|
||||
|
||||
if not state:
|
||||
@@ -422,9 +522,8 @@ def prepare_simulation():
|
||||
"error": t('api.simulationNotFound', id=simulation_id)
|
||||
}), 404
|
||||
|
||||
# 检查是否强制重新生成
|
||||
# Check whether a forced regeneration was requested.
|
||||
force_regenerate = data.get('force_regenerate', False)
|
||||
logger.info(f"开始处理 /prepare 请求: simulation_id={simulation_id}, force_regenerate={force_regenerate}")
|
||||
|
||||
# 检查是否已经准备完成(避免重复生成)
|
||||
if not force_regenerate:
|
||||
@@ -473,13 +572,18 @@ def prepare_simulation():
|
||||
# 这样前端在调用prepare后立即就能获取到预期Agent总数
|
||||
try:
|
||||
logger.info(f"同步获取实体数量: graph_id={state.graph_id}")
|
||||
reader = ZepEntityReader()
|
||||
# 快速读取实体(不需要边信息,只统计数量)
|
||||
filtered_preview = reader.filter_defined_entities(
|
||||
graph_id=state.graph_id,
|
||||
defined_entity_types=entity_types_list,
|
||||
enrich_with_edges=False # 不获取边信息,加快速度
|
||||
)
|
||||
reader = manager.create_entity_reader(state.graph_id)
|
||||
try:
|
||||
# Quick read without edge enrichment for immediate progress counts.
|
||||
filtered_preview = reader.filter_defined_entities(
|
||||
graph_id=state.graph_id,
|
||||
defined_entity_types=entity_types_list,
|
||||
enrich_with_edges=False,
|
||||
)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
# 保存实体数量到状态(供前端立即获取)
|
||||
state.entities_count = filtered_preview.filtered_count
|
||||
state.entity_types = list(filtered_preview.entity_types)
|
||||
@@ -490,11 +594,15 @@ def prepare_simulation():
|
||||
|
||||
# 创建异步任务
|
||||
task_manager = TaskManager()
|
||||
actor = current_actor()
|
||||
task_id = task_manager.create_task(
|
||||
task_type="simulation_prepare",
|
||||
metadata={
|
||||
"simulation_id": simulation_id,
|
||||
"project_id": state.project_id
|
||||
"project_id": state.project_id,
|
||||
"graph_id": state.graph_id,
|
||||
"organization_id": actor.organization_id,
|
||||
"owner_user_id": actor.user_id,
|
||||
}
|
||||
)
|
||||
|
||||
@@ -598,14 +706,14 @@ def prepare_simulation():
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"准备模拟失败: {str(e)}")
|
||||
task_manager.fail_task(task_id, str(e))
|
||||
logger.error("准备模拟失败: error_type=%s", type(e).__name__)
|
||||
task_manager.fail_task(task_id, t('api.internalError'))
|
||||
|
||||
# 更新模拟状态为失败
|
||||
state = manager.get_simulation(simulation_id)
|
||||
if state:
|
||||
state.status = SimulationStatus.FAILED
|
||||
state.error = str(e)
|
||||
state.error = t('api.internalError')
|
||||
manager._save_simulation_state(state)
|
||||
|
||||
# 启动后台线程
|
||||
@@ -625,19 +733,11 @@ def prepare_simulation():
|
||||
}
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 404
|
||||
except ValueError:
|
||||
return _safe_client_failure(404)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"启动准备任务失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/prepare/status', methods=['POST'])
|
||||
@@ -746,11 +846,7 @@ def get_prepare_status():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"查询任务状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 500
|
||||
return _safe_internal_failure("prepare task status", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>', methods=['GET'])
|
||||
@@ -758,13 +854,7 @@ def get_simulation(simulation_id: str):
|
||||
"""获取模拟状态"""
|
||||
try:
|
||||
manager = SimulationManager()
|
||||
state = manager.get_simulation(simulation_id)
|
||||
|
||||
if not state:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.simulationNotFound', id=simulation_id)
|
||||
}), 404
|
||||
state = require_scoped_simulation(simulation_id)
|
||||
|
||||
result = state.to_dict()
|
||||
|
||||
@@ -778,12 +868,7 @@ def get_simulation(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取模拟状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/list', methods=['GET'])
|
||||
@@ -797,8 +882,7 @@ def list_simulations():
|
||||
try:
|
||||
project_id = request.args.get('project_id')
|
||||
|
||||
manager = SimulationManager()
|
||||
simulations = manager.list_simulations(project_id=project_id)
|
||||
simulations = scoped_simulations(project_id=project_id)
|
||||
|
||||
return jsonify({
|
||||
"success": True,
|
||||
@@ -807,12 +891,7 @@ def list_simulations():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"列出模拟失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
def _get_report_id_for_simulation(simulation_id: str) -> str:
|
||||
@@ -913,7 +992,7 @@ def get_simulation_history():
|
||||
limit = request.args.get('limit', 20, type=int)
|
||||
|
||||
manager = SimulationManager()
|
||||
simulations = manager.list_simulations()[:limit]
|
||||
simulations = scoped_simulations()[:limit]
|
||||
|
||||
# 增强模拟数据,只从 Simulation 文件读取
|
||||
enriched_simulations = []
|
||||
@@ -980,12 +1059,7 @@ def get_simulation_history():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取历史模拟失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/profiles', methods=['GET'])
|
||||
@@ -1011,19 +1085,11 @@ def get_simulation_profiles(simulation_id: str):
|
||||
}
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 404
|
||||
except ValueError:
|
||||
return _safe_client_failure(404)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取Profile失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/profiles/realtime', methods=['GET'])
|
||||
@@ -1128,12 +1194,7 @@ def get_simulation_profiles_realtime(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"实时获取Profile失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/config/realtime', methods=['GET'])
|
||||
@@ -1248,12 +1309,7 @@ def get_simulation_config_realtime(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"实时获取Config失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/config', methods=['GET'])
|
||||
@@ -1284,12 +1340,7 @@ def get_simulation_config(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取配置失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/config/download', methods=['GET'])
|
||||
@@ -1313,12 +1364,7 @@ def download_simulation_config(simulation_id: str):
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"下载配置失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/script/<script_name>/download', methods=['GET'])
|
||||
@@ -1365,17 +1411,13 @@ def download_simulation_script(script_name: str):
|
||||
)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"下载脚本失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== Profile生成接口(独立使用) ==============
|
||||
|
||||
@simulation_bp.route('/generate-profiles', methods=['POST'])
|
||||
@idempotent
|
||||
def generate_profiles():
|
||||
"""
|
||||
直接从图谱生成OASIS Agent Profile(不创建模拟)
|
||||
@@ -1402,24 +1444,40 @@ def generate_profiles():
|
||||
use_llm = data.get('use_llm', True)
|
||||
platform = data.get('platform', 'reddit')
|
||||
|
||||
reader = ZepEntityReader()
|
||||
filtered = reader.filter_defined_entities(
|
||||
graph_id=graph_id,
|
||||
defined_entity_types=entity_types,
|
||||
enrich_with_edges=True
|
||||
)
|
||||
|
||||
if filtered.filtered_count == 0:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.noMatchingEntities')
|
||||
}), 400
|
||||
|
||||
generator = OasisProfileGenerator()
|
||||
profiles = generator.generate_profiles_from_entities(
|
||||
entities=filtered.entities,
|
||||
use_llm=use_llm
|
||||
)
|
||||
reader = _entity_reader_for_request(graph_id)
|
||||
try:
|
||||
filtered = reader.filter_defined_entities(
|
||||
graph_id=graph_id,
|
||||
defined_entity_types=entity_types,
|
||||
enrich_with_edges=True
|
||||
)
|
||||
|
||||
if filtered.filtered_count == 0:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.noMatchingEntities')
|
||||
}), 400
|
||||
|
||||
local_memory_tools = None
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
repository = getattr(reader, "repository", None)
|
||||
if repository is None:
|
||||
raise ApiError("local_memory_reader_required", 500, "api.internalError")
|
||||
local_memory_tools = LocalMemoryTools(repository)
|
||||
|
||||
generator = OasisProfileGenerator(
|
||||
graph_id=graph_id,
|
||||
use_zep_context=Config.MEMORY_BACKEND != "local",
|
||||
local_memory_tools=local_memory_tools,
|
||||
)
|
||||
profiles = generator.generate_profiles_from_entities(
|
||||
entities=filtered.entities,
|
||||
use_llm=use_llm
|
||||
)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
|
||||
if platform == "reddit":
|
||||
profiles_data = [p.to_reddit_format() for p in profiles]
|
||||
@@ -1439,12 +1497,7 @@ def generate_profiles():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"生成Profile失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== 模拟运行控制接口 ==============
|
||||
@@ -1485,6 +1538,7 @@ def _filter_simulation_agents(simulation_id: str, selected_agent_ids: list):
|
||||
logger.info(f"Twitter profiles: {len(rows)} -> {len(filtered)}")
|
||||
|
||||
@simulation_bp.route('/start', methods=['POST'])
|
||||
@idempotent
|
||||
def start_simulation():
|
||||
"""
|
||||
开始运行模拟
|
||||
@@ -1592,7 +1646,7 @@ def start_simulation():
|
||||
try:
|
||||
SimulationRunner.stop_simulation(simulation_id)
|
||||
except Exception as e:
|
||||
logger.warning(f"停止模拟时出现警告: {str(e)}")
|
||||
logger.warning("停止模拟时出现警告: error_type=%s", type(e).__name__)
|
||||
else:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
@@ -1648,7 +1702,9 @@ def start_simulation():
|
||||
platform=platform,
|
||||
max_rounds=max_rounds,
|
||||
enable_graph_memory_update=enable_graph_memory_update,
|
||||
graph_id=graph_id
|
||||
graph_id=graph_id,
|
||||
organization_id=current_actor().organization_id if Config.MEMORY_BACKEND == "local" else None,
|
||||
session_factory=current_app.extensions.get("crowdsight_session_factory") if Config.MEMORY_BACKEND == "local" else None,
|
||||
)
|
||||
|
||||
# 更新模拟状态
|
||||
@@ -1668,19 +1724,11 @@ def start_simulation():
|
||||
"data": response_data
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"启动模拟失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/stop', methods=['POST'])
|
||||
@@ -1727,19 +1775,11 @@ def stop_simulation():
|
||||
"data": run_state.to_dict()
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"停止模拟失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== 实时状态监控接口 ==============
|
||||
@@ -1794,12 +1834,7 @@ def get_run_status(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取运行状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/run-status/detail', methods=['GET'])
|
||||
@@ -1895,12 +1930,7 @@ def get_run_status_detail(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取详细状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/actions', methods=['GET'])
|
||||
@@ -1949,12 +1979,7 @@ def get_simulation_actions(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取动作历史失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/timeline', methods=['GET'])
|
||||
@@ -1989,12 +2014,7 @@ def get_simulation_timeline(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取时间线失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/agent-stats', methods=['GET'])
|
||||
@@ -2016,12 +2036,7 @@ def get_agent_stats(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取Agent统计失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== 数据库查询接口 ==============
|
||||
@@ -2096,12 +2111,7 @@ def get_simulation_posts(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取帖子失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/<simulation_id>/comments', methods=['GET'])
|
||||
@@ -2171,12 +2181,7 @@ def get_simulation_comments(simulation_id: str):
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取评论失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
# ============== Interview 采访接口 ==============
|
||||
@@ -2289,25 +2294,14 @@ def interview_agent():
|
||||
"data": result
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except TimeoutError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.interviewTimeout', error=str(e))
|
||||
}), 504
|
||||
except TimeoutError:
|
||||
return _safe_client_failure(504, "interview_timeout")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Interview失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/interview/batch', methods=['POST'])
|
||||
@@ -2427,25 +2421,14 @@ def interview_agents_batch():
|
||||
"data": result
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except TimeoutError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.batchInterviewTimeout', error=str(e))
|
||||
}), 504
|
||||
except TimeoutError:
|
||||
return _safe_client_failure(504, "batch_interview_timeout")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"批量Interview失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/interview/all', methods=['POST'])
|
||||
@@ -2530,25 +2513,14 @@ def interview_all_agents():
|
||||
"data": result
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except TimeoutError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": t('api.globalInterviewTimeout', error=str(e))
|
||||
}), 504
|
||||
except TimeoutError:
|
||||
return _safe_client_failure(504, "global_interview_timeout")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"全局Interview失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/interview/history', methods=['POST'])
|
||||
@@ -2615,12 +2587,7 @@ def get_interview_history():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取Interview历史失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/env-status', methods=['POST'])
|
||||
@@ -2680,12 +2647,7 @@ def get_env_status():
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"获取环境状态失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
|
||||
@simulation_bp.route('/close-env', methods=['POST'])
|
||||
@@ -2743,16 +2705,8 @@ def close_simulation_env():
|
||||
"data": result
|
||||
})
|
||||
|
||||
except ValueError as e:
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e)
|
||||
}), 400
|
||||
except ValueError:
|
||||
return _safe_client_failure(400)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"关闭环境失败: {str(e)}")
|
||||
return jsonify({
|
||||
"success": False,
|
||||
"error": str(e),
|
||||
"traceback": traceback.format_exc()
|
||||
}), 500
|
||||
return _safe_internal_failure("simulation operation", e)
|
||||
|
||||
@@ -9,6 +9,9 @@ from flask import Blueprint, request, jsonify
|
||||
from ..utils.llm_client import LLMClient
|
||||
from ..utils.locale import t, get_locale, get_language_instruction
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.api_errors import internal_error_payload
|
||||
from ..security.auth import require_auth
|
||||
from ..services.idempotency import idempotent
|
||||
|
||||
logger = get_logger('crowdsight.template')
|
||||
|
||||
@@ -22,6 +25,7 @@ with open(_templates_path, 'r', encoding='utf-8') as f:
|
||||
|
||||
|
||||
@template_bp.route('/list', methods=['GET'])
|
||||
@require_auth
|
||||
def list_templates():
|
||||
"""Return all available templates"""
|
||||
locale = get_locale()
|
||||
@@ -42,6 +46,8 @@ def list_templates():
|
||||
|
||||
|
||||
@template_bp.route('/auto-select', methods=['POST'])
|
||||
@require_auth
|
||||
@idempotent
|
||||
def auto_select_template():
|
||||
"""
|
||||
Analyze seed data and recommend the best template + pre-fill prompt.
|
||||
@@ -129,11 +135,12 @@ Return JSON:
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"Template auto-select failed: {e}")
|
||||
return jsonify({'success': False, 'error': str(e)}), 500
|
||||
logger.error("Template auto-select failed: error_type=%s", type(e).__name__)
|
||||
return jsonify(internal_error_payload(t)), 500
|
||||
|
||||
|
||||
@template_bp.route('/<template_id>/filter-rules', methods=['GET'])
|
||||
@require_auth
|
||||
def get_filter_rules(template_id):
|
||||
"""Return entity filter rules for a template (used by Step 1)"""
|
||||
for tmpl in _templates:
|
||||
|
||||
@@ -61,9 +61,16 @@ _llm_base_url, _llm_model_name, _llm_provider = _resolve_llm_config()
|
||||
class Config:
|
||||
"""Flask配置类"""
|
||||
|
||||
# Flask配置
|
||||
SECRET_KEY = os.environ.get('SECRET_KEY', 'crowdsight-secret-key')
|
||||
DEBUG = os.environ.get('FLASK_DEBUG', 'True').lower() == 'true'
|
||||
# Flask configuration: production-safe defaults; secrets must be supplied.
|
||||
SECRET_KEY = os.environ.get('SECRET_KEY')
|
||||
DEBUG = os.environ.get('FLASK_DEBUG', 'False').lower() == 'true'
|
||||
# Session cookies must be Secure in production; local HTTP can opt out explicitly.
|
||||
SESSION_COOKIE_SECURE = os.environ.get('SESSION_COOKIE_SECURE', 'False').lower() == 'true'
|
||||
CORS_ALLOWED_ORIGINS = [
|
||||
origin.strip()
|
||||
for origin in os.environ.get('CORS_ALLOWED_ORIGINS', 'http://localhost:3000').split(',')
|
||||
if origin.strip()
|
||||
]
|
||||
|
||||
# JSON配置 - 禁用ASCII转义,让中文直接显示(而不是 \uXXXX 格式)
|
||||
JSON_AS_ASCII = False
|
||||
@@ -74,7 +81,8 @@ class Config:
|
||||
LLM_BASE_URL = _llm_base_url
|
||||
LLM_MODEL_NAME = _llm_model_name
|
||||
|
||||
# Zep配置
|
||||
# Memory backend migration switch: keep Zep as the explicit default until data parity/cutover is complete.
|
||||
MEMORY_BACKEND = os.environ.get('MEMORY_BACKEND', 'zep').strip().lower()
|
||||
ZEP_API_KEY = os.environ.get('ZEP_API_KEY')
|
||||
|
||||
# 文件上传配置
|
||||
@@ -109,10 +117,14 @@ class Config:
|
||||
def validate(cls) -> list[str]:
|
||||
"""验证必要配置"""
|
||||
errors: list[str] = []
|
||||
if not cls.SECRET_KEY:
|
||||
errors.append("SECRET_KEY not configured")
|
||||
if not cls.LLM_API_KEY:
|
||||
errors.append("LLM_API_KEY 未配置")
|
||||
if not cls.ZEP_API_KEY:
|
||||
errors.append("ZEP_API_KEY 未配置")
|
||||
errors.append("LLM_API_KEY not configured")
|
||||
if cls.MEMORY_BACKEND not in {"zep", "local"}:
|
||||
errors.append("MEMORY_BACKEND must be zep or local")
|
||||
elif cls.MEMORY_BACKEND == "zep" and not cls.ZEP_API_KEY:
|
||||
errors.append("ZEP_API_KEY not configured")
|
||||
return errors
|
||||
|
||||
@classmethod
|
||||
|
||||
50
backend/app/db.py
Normal file
50
backend/app/db.py
Normal file
@@ -0,0 +1,50 @@
|
||||
"""Database engine and declarative base helpers."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from pathlib import Path
|
||||
|
||||
from sqlalchemy import create_engine, event
|
||||
from sqlalchemy.engine import Engine
|
||||
from sqlalchemy.orm import DeclarativeBase, Session, sessionmaker
|
||||
|
||||
|
||||
class Base(DeclarativeBase):
|
||||
pass
|
||||
|
||||
|
||||
def create_database_engine(database_url: str | None = None, **kwargs) -> Engine:
|
||||
"""Create a configured SQLAlchemy engine without opening a global session."""
|
||||
url = database_url or os.environ.get("DATABASE_URL")
|
||||
if not url:
|
||||
data_dir = Path(os.environ.get("CROWDSIGHT_DATA_DIR", "backend/uploads"))
|
||||
data_dir.mkdir(parents=True, exist_ok=True)
|
||||
url = f"sqlite+pysqlite:///{(data_dir / 'crowdsight.db').resolve()}"
|
||||
|
||||
connect_args = dict(kwargs.pop("connect_args", {}))
|
||||
if url.startswith("sqlite"):
|
||||
connect_args.setdefault("check_same_thread", False)
|
||||
|
||||
engine = create_engine(
|
||||
url,
|
||||
future=True,
|
||||
pool_pre_ping=True,
|
||||
connect_args=connect_args,
|
||||
**kwargs,
|
||||
)
|
||||
if url.startswith("sqlite"):
|
||||
@event.listens_for(engine, "connect")
|
||||
def _enable_sqlite_foreign_keys(dbapi_connection, _connection_record):
|
||||
cursor = dbapi_connection.cursor()
|
||||
try:
|
||||
cursor.execute("PRAGMA foreign_keys=ON")
|
||||
finally:
|
||||
cursor.close()
|
||||
|
||||
return engine
|
||||
|
||||
|
||||
def create_session_factory(engine: Engine) -> sessionmaker[Session]:
|
||||
"""Return a factory; callers own transaction boundaries and commits."""
|
||||
return sessionmaker(bind=engine, autoflush=True, expire_on_commit=False)
|
||||
@@ -4,6 +4,21 @@
|
||||
|
||||
from .task import TaskManager, TaskStatus
|
||||
from .project import Project, ProjectStatus, ProjectManager
|
||||
from .saas import AuthSession, Membership, Organization, User
|
||||
from .memory import MemoryEdge, MemoryEpisode, MemoryGraph, MemoryNode
|
||||
from .operations import AuditLog, IdempotencyRecord, Job, JobStatus
|
||||
from .product import DurableReport, ProductProject, ProductSimulation, ProjectStatus, ReportStatus, SimulationStatus
|
||||
from .settings import PlatformSettings
|
||||
from .rate_limit import RateLimitEvent
|
||||
from .usage import UsageEvent
|
||||
from .password_reset import PasswordResetToken
|
||||
|
||||
__all__ = ['TaskManager', 'TaskStatus', 'Project', 'ProjectStatus', 'ProjectManager']
|
||||
__all__ = [
|
||||
'TaskManager', 'TaskStatus', 'Project', 'ProjectStatus', 'ProjectManager',
|
||||
'AuthSession', 'Membership', 'Organization', 'User',
|
||||
'MemoryEdge', 'MemoryEpisode', 'MemoryGraph', 'MemoryNode',
|
||||
'AuditLog', 'IdempotencyRecord', 'Job', 'JobStatus',
|
||||
'ProductProject', 'ProductSimulation', 'DurableReport', 'SimulationStatus', 'ReportStatus',
|
||||
'PlatformSettings', 'RateLimitEvent', 'UsageEvent', 'PasswordResetToken',
|
||||
]
|
||||
|
||||
|
||||
141
backend/app/models/memory.py
Normal file
141
backend/app/models/memory.py
Normal file
@@ -0,0 +1,141 @@
|
||||
"""Durable local graph-memory schema replacing the storage side of Zep."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import DateTime, Float, ForeignKey, Index, JSON, String, Text, UniqueConstraint, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _id(prefix: str) -> str:
|
||||
return f"{prefix}_{uuid4().hex}"
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class MemoryGraph(Base):
|
||||
__tablename__ = "memory_graphs"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(128), primary_key=True, default=lambda: _id("graph"))
|
||||
organization_id: Mapped[str] = mapped_column(String(64), nullable=False, index=True)
|
||||
project_id: Mapped[str] = mapped_column(String(128), nullable=False, index=True)
|
||||
ontology: Mapped[dict] = mapped_column(JSON, nullable=False, default=dict, server_default=text("'{}'"))
|
||||
status: Mapped[str] = mapped_column(String(32), nullable=False, default="ready", server_default=text("'ready'"))
|
||||
version: Mapped[int] = mapped_column(nullable=False, default=1, server_default=text("1"))
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now(), onupdate=_utc_now
|
||||
)
|
||||
|
||||
episodes: Mapped[list["MemoryEpisode"]] = relationship(
|
||||
back_populates="graph", cascade="all, delete-orphan"
|
||||
)
|
||||
nodes: Mapped[list["MemoryNode"]] = relationship(
|
||||
back_populates="graph", cascade="all, delete-orphan"
|
||||
)
|
||||
edges: Mapped[list["MemoryEdge"]] = relationship(
|
||||
back_populates="graph", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
|
||||
class MemoryEpisode(Base):
|
||||
__tablename__ = "memory_episodes"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("graph_id", "source_type", "source_ref", name="uq_memory_episode_source"),
|
||||
Index("ix_memory_episodes_graph_status", "graph_id", "status"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(128), primary_key=True, default=lambda: _id("episode"))
|
||||
graph_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memory_graphs.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
source_type: Mapped[str] = mapped_column(String(32), nullable=False)
|
||||
source_ref: Mapped[str] = mapped_column(String(256), nullable=False)
|
||||
normalized_text: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
summary: Mapped[str] = mapped_column(Text, nullable=False, default="", server_default=text("''"))
|
||||
status: Mapped[str] = mapped_column(String(32), nullable=False, default="processed", server_default=text("'processed'"))
|
||||
extractor_version: Mapped[str] = mapped_column(String(64), nullable=False, default="v1", server_default=text("'v1'"))
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
|
||||
graph: Mapped[MemoryGraph] = relationship(back_populates="episodes")
|
||||
|
||||
|
||||
class MemoryNode(Base):
|
||||
__tablename__ = "memory_nodes"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("graph_id", "normalized_name", name="uq_memory_node_graph_name"),
|
||||
Index("ix_memory_nodes_graph_name", "graph_id", "normalized_name"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(128), primary_key=True, default=lambda: _id("node"))
|
||||
graph_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memory_graphs.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
canonical_name: Mapped[str] = mapped_column(String(512), nullable=False)
|
||||
normalized_name: Mapped[str] = mapped_column(String(512), nullable=False)
|
||||
labels: Mapped[list] = mapped_column(JSON, nullable=False, default=list, server_default=text("'[]'"))
|
||||
aliases: Mapped[list] = mapped_column(JSON, nullable=False, default=list, server_default=text("'[]'"))
|
||||
attributes: Mapped[dict] = mapped_column(JSON, nullable=False, default=dict, server_default=text("'{}'"))
|
||||
summary: Mapped[str] = mapped_column(Text, nullable=False, default="", server_default=text("''"))
|
||||
confidence: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, server_default=text("0"))
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now(), onupdate=_utc_now
|
||||
)
|
||||
|
||||
graph: Mapped[MemoryGraph] = relationship(back_populates="nodes")
|
||||
outgoing_edges: Mapped[list["MemoryEdge"]] = relationship(
|
||||
foreign_keys="MemoryEdge.source_node_id", back_populates="source_node"
|
||||
)
|
||||
incoming_edges: Mapped[list["MemoryEdge"]] = relationship(
|
||||
foreign_keys="MemoryEdge.target_node_id", back_populates="target_node"
|
||||
)
|
||||
|
||||
|
||||
class MemoryEdge(Base):
|
||||
__tablename__ = "memory_edges"
|
||||
__table_args__ = (
|
||||
Index("ix_memory_edges_graph_relation", "graph_id", "relation"),
|
||||
Index("ix_memory_edges_graph_temporal", "graph_id", "valid_at", "invalid_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(128), primary_key=True, default=lambda: _id("edge"))
|
||||
graph_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memory_graphs.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
source_node_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memory_nodes.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
target_node_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memory_nodes.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
relation: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
fact: Mapped[str] = mapped_column(Text, nullable=False)
|
||||
attributes: Mapped[dict] = mapped_column(JSON, nullable=False, default=dict, server_default=text("'{}'"))
|
||||
confidence: Mapped[float] = mapped_column(Float, nullable=False, default=0.0, server_default=text("0"))
|
||||
valid_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
invalid_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
expired_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
|
||||
graph: Mapped[MemoryGraph] = relationship(back_populates="edges")
|
||||
source_node: Mapped[MemoryNode] = relationship(
|
||||
foreign_keys=[source_node_id], back_populates="outgoing_edges"
|
||||
)
|
||||
target_node: Mapped[MemoryNode] = relationship(
|
||||
foreign_keys=[target_node_id], back_populates="incoming_edges"
|
||||
)
|
||||
128
backend/app/models/operations.py
Normal file
128
backend/app/models/operations.py
Normal file
@@ -0,0 +1,128 @@
|
||||
"""Durable operation metadata for jobs, retries, idempotency, and audit."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import JSON, CheckConstraint, DateTime, ForeignKey, Index, Integer, String, Text, UniqueConstraint, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _id(prefix: str) -> str:
|
||||
return f"{prefix}_{uuid4().hex}"
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class JobStatus(str, Enum):
|
||||
QUEUED = "queued"
|
||||
RUNNING = "running"
|
||||
SUCCEEDED = "succeeded"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class Job(Base):
|
||||
"""Durable, tenant-owned unit of asynchronous work."""
|
||||
|
||||
__tablename__ = "jobs"
|
||||
__table_args__ = (
|
||||
CheckConstraint(
|
||||
"status IN ('queued', 'running', 'succeeded', 'failed', 'cancelled')",
|
||||
name="ck_jobs_status",
|
||||
),
|
||||
Index("ix_jobs_org_status_created", "organization_id", "status", "created_at"),
|
||||
Index("ix_jobs_org_owner", "organization_id", "owner_user_id"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("job"))
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
owner_user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True, index=True
|
||||
)
|
||||
project_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True)
|
||||
graph_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True)
|
||||
operation: Mapped[str] = mapped_column(String(120), nullable=False)
|
||||
status: Mapped[JobStatus] = mapped_column(
|
||||
String(32), nullable=False, default=JobStatus.QUEUED.value, server_default=text("'queued'")
|
||||
)
|
||||
progress: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default=text("0"))
|
||||
message: Mapped[str] = mapped_column(Text, nullable=False, default="", server_default=text("''"))
|
||||
result: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
progress_detail: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
job_metadata: Mapped[dict | list | None] = mapped_column("metadata", JSON, nullable=True)
|
||||
error_code: Mapped[str | None] = mapped_column(String(120), nullable=True)
|
||||
result_ref: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
idempotency_key: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
attempt: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default=text("0"))
|
||||
settings_version: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, onupdate=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class IdempotencyRecord(Base):
|
||||
"""Request fingerprint and replayable response for retry-safe mutations."""
|
||||
|
||||
__tablename__ = "idempotency_records"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("organization_id", "user_id", "key", name="uq_idempotency_org_user_key"),
|
||||
Index("ix_idempotency_expiry", "expires_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("idem"))
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
key: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
request_hash: Mapped[str] = mapped_column(String(64), nullable=False)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="reserved", server_default=text("'reserved'")
|
||||
)
|
||||
response_status: Mapped[int | None] = mapped_column(Integer, nullable=True)
|
||||
response_body: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
completed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class AuditLog(Base):
|
||||
"""Tenant-scoped redacted audit event; details must never contain secrets."""
|
||||
|
||||
__tablename__ = "audit_logs"
|
||||
__table_args__ = (
|
||||
Index("ix_audit_org_created", "organization_id", "created_at"),
|
||||
Index("ix_audit_org_target", "organization_id", "target_type", "target_id"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("audit"))
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
actor_user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True, index=True
|
||||
)
|
||||
action: Mapped[str] = mapped_column(String(160), nullable=False)
|
||||
target_type: Mapped[str] = mapped_column(String(80), nullable=False)
|
||||
target_id: Mapped[str | None] = mapped_column(String(160), nullable=True)
|
||||
details: Mapped[dict | list | None] = mapped_column("metadata", JSON, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
37
backend/app/models/password_reset.py
Normal file
37
backend/app/models/password_reset.py
Normal file
@@ -0,0 +1,37 @@
|
||||
"""Durable, single-use password reset tokens."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, ForeignKey, String, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class PasswordResetToken(Base):
|
||||
"""One hash of a one-time reset token; plaintext is never stored."""
|
||||
|
||||
__tablename__ = "password_reset_tokens"
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(64), primary_key=True, default=lambda: f"prt_{uuid4().hex}"
|
||||
)
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
token_hash: Mapped[str] = mapped_column(String(128), nullable=False)
|
||||
auth_version: Mapped[int] = mapped_column(nullable=False, default=0, server_default=text("0"))
|
||||
used: Mapped[bool] = mapped_column(
|
||||
Boolean, nullable=False, default=False, server_default=text("0")
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
186
backend/app/models/product.py
Normal file
186
backend/app/models/product.py
Normal file
@@ -0,0 +1,186 @@
|
||||
"""Durable product-resource schema: projects, simulations, and reports.
|
||||
|
||||
These replace the legacy filesystem-backed ProjectManager / SimulationManager /
|
||||
ReportManager payloads with tenant- and owner-scoped SQL rows, so product state
|
||||
survives restarts and multiple workers without cross-tenant leakage.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from enum import Enum
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import (
|
||||
DateTime,
|
||||
ForeignKey,
|
||||
Index,
|
||||
Integer,
|
||||
JSON,
|
||||
String,
|
||||
Text,
|
||||
UniqueConstraint,
|
||||
func,
|
||||
text,
|
||||
)
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _id(prefix: str) -> str:
|
||||
return f"{prefix}_{uuid4().hex}"
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class ProjectStatus(str, Enum):
|
||||
CREATED = "created"
|
||||
ONTOLOGY_GENERATED = "ontology_generated"
|
||||
GRAPH_BUILDING = "graph_building"
|
||||
GRAPH_COMPLETED = "graph_completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class SimulationStatus(str, Enum):
|
||||
CREATED = "created"
|
||||
PREPARING = "preparing"
|
||||
READY = "ready"
|
||||
RUNNING = "running"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
CANCELLED = "cancelled"
|
||||
|
||||
|
||||
class ReportStatus(str, Enum):
|
||||
DRAFT = "draft"
|
||||
PLANNING = "planning"
|
||||
GENERATING = "generating"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
class ProductProject(Base):
|
||||
"""Durable, tenant-owned project record (metadata + ontology)."""
|
||||
|
||||
__tablename__ = "projects"
|
||||
__table_args__ = (
|
||||
Index("ix_projects_org_owner", "organization_id", "owner_user_id"),
|
||||
Index("ix_projects_org_created", "organization_id", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(128), primary_key=True, default=lambda: _id("project")
|
||||
)
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
owner_user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True, index=True
|
||||
)
|
||||
name: Mapped[str] = mapped_column(String(255), nullable=False, server_default=text("''"))
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default=ProjectStatus.CREATED.value, server_default=text("'created'")
|
||||
)
|
||||
language: Mapped[str] = mapped_column(
|
||||
String(16), nullable=False, default="en", server_default=text("'en'")
|
||||
)
|
||||
total_text_length: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, default=0, server_default=text("0")
|
||||
)
|
||||
source_metadata: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
ontology: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
analysis_summary: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
simulation_requirement: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
graph_id: Mapped[str | None] = mapped_column(String(128), nullable=True, index=True)
|
||||
graph_build_task_id: Mapped[str | None] = mapped_column(String(128), nullable=True)
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now(), onupdate=_utc_now
|
||||
)
|
||||
|
||||
|
||||
class ProductSimulation(Base):
|
||||
"""Durable, tenant-scoped simulation with a config snapshot."""
|
||||
|
||||
__tablename__ = "simulations"
|
||||
__table_args__ = (
|
||||
Index("ix_simulations_org_project", "organization_id", "project_id"),
|
||||
Index("ix_simulations_org_created", "organization_id", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(128), primary_key=True, default=lambda: _id("sim")
|
||||
)
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
project_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("projects.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
created_by_user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True
|
||||
)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default=SimulationStatus.CREATED.value, server_default=text("'created'")
|
||||
)
|
||||
platform: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="parallel", server_default=text("'parallel'")
|
||||
)
|
||||
config: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
current_round: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default=text("0"))
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now(), onupdate=_utc_now
|
||||
)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
|
||||
class DurableReport(Base):
|
||||
"""Durable, tenant-scoped report with outline/status/content metadata."""
|
||||
|
||||
__tablename__ = "reports"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("organization_id", "id", name="uq_reports_org_id"),
|
||||
Index("ix_reports_org_project", "organization_id", "project_id"),
|
||||
Index("ix_reports_org_simulation", "organization_id", "simulation_id"),
|
||||
Index("ix_reports_org_created", "organization_id", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(128), primary_key=True, default=lambda: _id("report")
|
||||
)
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
project_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("projects.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
simulation_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("simulations.id", ondelete="SET NULL"), nullable=True, index=True
|
||||
)
|
||||
created_by_user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True
|
||||
)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default=ReportStatus.DRAFT.value, server_default=text("'draft'")
|
||||
)
|
||||
title: Mapped[str] = mapped_column(String(255), nullable=False, server_default=text("''"))
|
||||
outline: Mapped[dict | list | None] = mapped_column(JSON, nullable=True)
|
||||
markdown_content: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
error: Mapped[str | None] = mapped_column(Text, nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
updated_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now(), onupdate=_utc_now
|
||||
)
|
||||
finished_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
@@ -5,8 +5,10 @@
|
||||
|
||||
import os
|
||||
import json
|
||||
import re
|
||||
import uuid
|
||||
import shutil
|
||||
import tempfile
|
||||
from datetime import datetime
|
||||
from typing import Dict, Any, List, Optional
|
||||
from enum import Enum
|
||||
@@ -31,6 +33,8 @@ class Project:
|
||||
status: ProjectStatus
|
||||
created_at: str
|
||||
updated_at: str
|
||||
organization_id: Optional[str] = None
|
||||
owner_user_id: Optional[str] = None
|
||||
|
||||
# 文件信息
|
||||
files: List[Dict[str, str]] = field(default_factory=list) # [{filename, path, size}]
|
||||
@@ -60,6 +64,8 @@ class Project:
|
||||
"status": self.status.value if isinstance(self.status, ProjectStatus) else self.status,
|
||||
"created_at": self.created_at,
|
||||
"updated_at": self.updated_at,
|
||||
"organization_id": self.organization_id,
|
||||
"owner_user_id": self.owner_user_id,
|
||||
"files": self.files,
|
||||
"total_text_length": self.total_text_length,
|
||||
"ontology": self.ontology,
|
||||
@@ -85,6 +91,8 @@ class Project:
|
||||
status=status,
|
||||
created_at=data.get('created_at', ''),
|
||||
updated_at=data.get('updated_at', ''),
|
||||
organization_id=data.get('organization_id'),
|
||||
owner_user_id=data.get('owner_user_id'),
|
||||
files=data.get('files', []),
|
||||
total_text_length=data.get('total_text_length', 0),
|
||||
ontology=data.get('ontology'),
|
||||
@@ -100,6 +108,8 @@ class Project:
|
||||
|
||||
class ProjectManager:
|
||||
"""项目管理器 - 负责项目的持久化存储和检索"""
|
||||
|
||||
_SAFE_PROJECT_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,127}$")
|
||||
|
||||
# 项目存储根目录
|
||||
PROJECTS_DIR = os.path.join(Config.UPLOAD_FOLDER, 'projects')
|
||||
@@ -109,10 +119,21 @@ class ProjectManager:
|
||||
"""确保项目目录存在"""
|
||||
os.makedirs(cls.PROJECTS_DIR, exist_ok=True)
|
||||
|
||||
@classmethod
|
||||
def _validate_project_id(cls, project_id: str) -> str:
|
||||
if not isinstance(project_id, str) or not cls._SAFE_PROJECT_ID.fullmatch(project_id):
|
||||
raise ValueError("invalid_project_id")
|
||||
return project_id
|
||||
|
||||
@classmethod
|
||||
def _get_project_dir(cls, project_id: str) -> str:
|
||||
"""获取项目目录路径"""
|
||||
return os.path.join(cls.PROJECTS_DIR, project_id)
|
||||
"""获取项目目录路径,拒绝路径分隔符和 traversal。"""
|
||||
safe_project_id = cls._validate_project_id(project_id)
|
||||
root = os.path.realpath(cls.PROJECTS_DIR)
|
||||
project_dir = os.path.realpath(os.path.join(root, safe_project_id))
|
||||
if os.path.commonpath([root, project_dir]) != root:
|
||||
raise ValueError("invalid_project_id")
|
||||
return project_dir
|
||||
|
||||
@classmethod
|
||||
def _get_project_meta_path(cls, project_id: str) -> str:
|
||||
@@ -130,7 +151,12 @@ class ProjectManager:
|
||||
return os.path.join(cls._get_project_dir(project_id), 'extracted_text.txt')
|
||||
|
||||
@classmethod
|
||||
def create_project(cls, name: str = "Unnamed Project") -> Project:
|
||||
def create_project(
|
||||
cls,
|
||||
name: str = "Unnamed Project",
|
||||
organization_id: Optional[str] = None,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> Project:
|
||||
"""
|
||||
创建新项目
|
||||
|
||||
@@ -150,7 +176,9 @@ class ProjectManager:
|
||||
name=name,
|
||||
status=ProjectStatus.CREATED,
|
||||
created_at=now,
|
||||
updated_at=now
|
||||
updated_at=now,
|
||||
organization_id=organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
|
||||
# 创建项目目录结构
|
||||
@@ -166,12 +194,24 @@ class ProjectManager:
|
||||
|
||||
@classmethod
|
||||
def save_project(cls, project: Project) -> None:
|
||||
"""保存项目元数据"""
|
||||
"""保存项目元数据,使用同目录临时文件+原子替换。"""
|
||||
project.updated_at = datetime.now().isoformat()
|
||||
project_dir = cls._get_project_dir(project.project_id)
|
||||
os.makedirs(project_dir, exist_ok=True)
|
||||
meta_path = cls._get_project_meta_path(project.project_id)
|
||||
|
||||
with open(meta_path, 'w', encoding='utf-8') as f:
|
||||
json.dump(project.to_dict(), f, ensure_ascii=False, indent=2)
|
||||
fd, temp_path = tempfile.mkstemp(prefix=".project-", suffix=".json", dir=project_dir)
|
||||
try:
|
||||
with os.fdopen(fd, 'w', encoding='utf-8') as f:
|
||||
json.dump(project.to_dict(), f, ensure_ascii=False, indent=2)
|
||||
f.flush()
|
||||
os.fsync(f.fileno())
|
||||
os.replace(temp_path, meta_path)
|
||||
except Exception:
|
||||
try:
|
||||
os.unlink(temp_path)
|
||||
except FileNotFoundError:
|
||||
pass
|
||||
raise
|
||||
|
||||
@classmethod
|
||||
def get_project(cls, project_id: str) -> Optional[Project]:
|
||||
@@ -184,18 +224,32 @@ class ProjectManager:
|
||||
Returns:
|
||||
Project对象,如果不存在返回None
|
||||
"""
|
||||
meta_path = cls._get_project_meta_path(project_id)
|
||||
|
||||
try:
|
||||
meta_path = cls._get_project_meta_path(project_id)
|
||||
except ValueError:
|
||||
return None
|
||||
|
||||
if not os.path.exists(meta_path):
|
||||
return None
|
||||
|
||||
with open(meta_path, 'r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
try:
|
||||
with open(meta_path, 'r', encoding='utf-8') as f:
|
||||
data = json.load(f)
|
||||
except (OSError, json.JSONDecodeError, TypeError, KeyError, ValueError):
|
||||
return None
|
||||
|
||||
return Project.from_dict(data)
|
||||
try:
|
||||
return Project.from_dict(data)
|
||||
except (KeyError, TypeError, ValueError):
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def list_projects(cls, limit: int = 50) -> List[Project]:
|
||||
def list_projects(
|
||||
cls,
|
||||
limit: int = 50,
|
||||
organization_id: Optional[str] = None,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> List[Project]:
|
||||
"""
|
||||
列出所有项目
|
||||
|
||||
@@ -210,30 +264,84 @@ class ProjectManager:
|
||||
projects = []
|
||||
for project_id in os.listdir(cls.PROJECTS_DIR):
|
||||
project = cls.get_project(project_id)
|
||||
if project:
|
||||
projects.append(project)
|
||||
if not project:
|
||||
continue
|
||||
if organization_id is not None and project.organization_id != organization_id:
|
||||
continue
|
||||
if owner_user_id is not None and project.owner_user_id != owner_user_id:
|
||||
continue
|
||||
projects.append(project)
|
||||
|
||||
# 按创建时间倒序排序
|
||||
projects.sort(key=lambda p: p.created_at, reverse=True)
|
||||
|
||||
return projects[:limit]
|
||||
|
||||
@classmethod
|
||||
def find_project_by_graph_id(
|
||||
cls,
|
||||
graph_id: str,
|
||||
*,
|
||||
organization_id: str,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> Optional[Project]:
|
||||
for project in cls.list_projects(
|
||||
organization_id=organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
):
|
||||
if project.graph_id == graph_id:
|
||||
return project
|
||||
return None
|
||||
|
||||
@classmethod
|
||||
def get_project_for_scope(
|
||||
cls,
|
||||
project_id: str,
|
||||
*,
|
||||
organization_id: str,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> Optional[Project]:
|
||||
"""Return only an explicitly owned/scoped project; legacy records fail closed."""
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
return None
|
||||
project = cls.get_project(project_id)
|
||||
if project is None or project.organization_id != organization_id:
|
||||
return None
|
||||
if owner_user_id is not None and project.owner_user_id != owner_user_id:
|
||||
return None
|
||||
return project
|
||||
|
||||
@classmethod
|
||||
def delete_project(cls, project_id: str) -> bool:
|
||||
def delete_project(
|
||||
cls,
|
||||
project_id: str,
|
||||
*,
|
||||
organization_id: str,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> bool:
|
||||
"""
|
||||
删除项目及其所有文件
|
||||
|
||||
删除项目及其所有文件,但只允许删除明确授权范围内的项目。
|
||||
|
||||
Args:
|
||||
project_id: 项目ID
|
||||
|
||||
organization_id: 当前请求的组织范围
|
||||
owner_user_id: 普通用户的所有者范围;管理员可留空
|
||||
|
||||
Returns:
|
||||
是否删除成功
|
||||
"""
|
||||
project_dir = cls._get_project_dir(project_id)
|
||||
|
||||
project = cls.get_project_for_scope(
|
||||
project_id,
|
||||
organization_id=organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
if project is None:
|
||||
return False
|
||||
|
||||
project_dir = cls._get_project_dir(project.project_id)
|
||||
if not os.path.exists(project_dir):
|
||||
return False
|
||||
|
||||
|
||||
shutil.rmtree(project_dir)
|
||||
return True
|
||||
|
||||
|
||||
35
backend/app/models/rate_limit.py
Normal file
35
backend/app/models/rate_limit.py
Normal file
@@ -0,0 +1,35 @@
|
||||
"""Durable rate-limit event records."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import DateTime, Index, String, func
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class RateLimitEvent(Base):
|
||||
"""One recorded rate-limit hit for an operation + key (no secrets)."""
|
||||
|
||||
__tablename__ = "rate_limit_events"
|
||||
__table_args__ = (
|
||||
Index("ix_rate_limit_op_key_created", "operation", "key", "created_at"),
|
||||
Index("ix_rate_limit_org_created", "organization_id", "created_at"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(64), primary_key=True, default=lambda: f"rl_{uuid4().hex}"
|
||||
)
|
||||
operation: Mapped[str] = mapped_column(String(120), nullable=False)
|
||||
key: Mapped[str] = mapped_column(String(255), nullable=False, index=True)
|
||||
organization_id: Mapped[str | None] = mapped_column(String(64), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
138
backend/app/models/saas.py
Normal file
138
backend/app/models/saas.py
Normal file
@@ -0,0 +1,138 @@
|
||||
"""SQLAlchemy identity and tenant metadata models."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import List
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import CheckConstraint, DateTime, Enum as SAEnum, ForeignKey, Integer, String, UniqueConstraint, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column, relationship
|
||||
|
||||
from ..db import Base
|
||||
from ..security.policy import Role
|
||||
|
||||
|
||||
def _id(prefix: str) -> str:
|
||||
return f"{prefix}_{uuid4().hex}"
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
def _role_values(enum_type):
|
||||
return [member.value for member in enum_type]
|
||||
|
||||
|
||||
class Organization(Base):
|
||||
__tablename__ = "organizations"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("org"))
|
||||
name: Mapped[str] = mapped_column(String(160), nullable=False)
|
||||
slug: Mapped[str] = mapped_column(String(80), nullable=False, unique=True, index=True)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="active", server_default=text("'active'")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
memberships: Mapped[List["Membership"]] = relationship(
|
||||
back_populates="organization", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
|
||||
class User(Base):
|
||||
__tablename__ = "users"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("usr"))
|
||||
email_normalized: Mapped[str] = mapped_column(String(320), nullable=False, unique=True, index=True)
|
||||
password_hash: Mapped[str] = mapped_column(
|
||||
String(512), nullable=False, default="!invite_pending", server_default=text("'!invite_pending'")
|
||||
)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="active", server_default=text("'active'")
|
||||
)
|
||||
auth_version: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, default=0, server_default=text("0")
|
||||
)
|
||||
locale: Mapped[str] = mapped_column(
|
||||
String(8), nullable=False, default="th", server_default=text("'th'")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
last_login_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
memberships: Mapped[List["Membership"]] = relationship(
|
||||
back_populates="user", cascade="all, delete-orphan"
|
||||
)
|
||||
sessions: Mapped[List["AuthSession"]] = relationship(
|
||||
back_populates="user", cascade="all, delete-orphan"
|
||||
)
|
||||
|
||||
|
||||
class Membership(Base):
|
||||
__tablename__ = "memberships"
|
||||
__table_args__ = (
|
||||
UniqueConstraint("user_id", "organization_id", name="uq_membership_user_org"),
|
||||
CheckConstraint(
|
||||
"role IN ('super_admin', 'admin', 'user')",
|
||||
name="ck_membership_role",
|
||||
),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("mem"))
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
role: Mapped[Role] = mapped_column(
|
||||
SAEnum(
|
||||
Role,
|
||||
name="role",
|
||||
values_callable=_role_values,
|
||||
native_enum=False,
|
||||
create_constraint=False,
|
||||
validate_strings=True,
|
||||
),
|
||||
nullable=False,
|
||||
)
|
||||
status: Mapped[str] = mapped_column(
|
||||
String(32), nullable=False, default="active", server_default=text("'active'")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
|
||||
user: Mapped[User] = relationship(back_populates="memberships")
|
||||
organization: Mapped[Organization] = relationship(back_populates="memberships")
|
||||
sessions: Mapped[List["AuthSession"]] = relationship(back_populates="membership")
|
||||
|
||||
|
||||
class AuthSession(Base):
|
||||
__tablename__ = "sessions"
|
||||
|
||||
id: Mapped[str] = mapped_column(String(64), primary_key=True, default=lambda: _id("ses"))
|
||||
user_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
membership_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("memberships.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
token_hash: Mapped[str] = mapped_column(String(64), nullable=False, unique=True, index=True)
|
||||
auth_version: Mapped[int] = mapped_column(
|
||||
Integer, nullable=False, default=0, server_default=text("0")
|
||||
)
|
||||
expires_at: Mapped[datetime] = mapped_column(DateTime(timezone=True), nullable=False)
|
||||
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), default=_utc_now, server_default=func.now(), nullable=False
|
||||
)
|
||||
last_seen_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
|
||||
|
||||
user: Mapped[User] = relationship(back_populates="sessions")
|
||||
membership: Mapped[Membership] = relationship(back_populates="sessions")
|
||||
42
backend/app/models/settings.py
Normal file
42
backend/app/models/settings.py
Normal file
@@ -0,0 +1,42 @@
|
||||
"""Durable, versioned platform settings (LLM provider etc.) with redacted secrets."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import Boolean, DateTime, ForeignKey, JSON, String, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class PlatformSettings(Base):
|
||||
"""A versioned snapshot of platform LLM settings.
|
||||
|
||||
Public/non-secret settings live in ``settings`` (JSON). The API key must be
|
||||
stored encrypted (as ``secret_ref``), never as plaintext in ``settings``.
|
||||
``active`` marks the current effective version.
|
||||
"""
|
||||
|
||||
__tablename__ = "platform_settings"
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(64), primary_key=True, default=lambda: f"ps_{uuid4().hex}"
|
||||
)
|
||||
version: Mapped[str] = mapped_column(String(64), nullable=False, unique=True)
|
||||
settings: Mapped[dict | None] = mapped_column(JSON, nullable=True)
|
||||
secret_ref: Mapped[str | None] = mapped_column(String(512), nullable=True)
|
||||
updated_by_user_id: Mapped[str | None] = mapped_column(
|
||||
String(64), nullable=True
|
||||
)
|
||||
active: Mapped[bool] = mapped_column(
|
||||
Boolean, nullable=False, default=False, server_default=text("0")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
@@ -1,43 +1,48 @@
|
||||
"""
|
||||
任务状态管理
|
||||
用于跟踪长时间运行的任务(如图谱构建)
|
||||
"""Durable task status management with a test-only in-memory fallback.
|
||||
|
||||
The Flask application configures a SQLAlchemy session factory at startup. Code
|
||||
that uses TaskManager outside an application (small unit tests and legacy
|
||||
adapters) keeps the old in-memory behavior, but production requests do not.
|
||||
"""
|
||||
|
||||
import uuid
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
from typing import Dict, Any, Optional
|
||||
import uuid
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from enum import Enum
|
||||
from typing import Any, Dict, Optional, cast
|
||||
|
||||
from flask import current_app, has_app_context
|
||||
from sqlalchemy import delete, select
|
||||
|
||||
from ..models.operations import Job, JobStatus
|
||||
from ..utils.locale import t
|
||||
|
||||
|
||||
class TaskStatus(str, Enum):
|
||||
"""任务状态枚举"""
|
||||
PENDING = "pending" # 等待中
|
||||
PROCESSING = "processing" # 处理中
|
||||
COMPLETED = "completed" # 已完成
|
||||
FAILED = "failed" # 失败
|
||||
PENDING = "pending"
|
||||
PROCESSING = "processing"
|
||||
COMPLETED = "completed"
|
||||
FAILED = "failed"
|
||||
|
||||
|
||||
@dataclass
|
||||
class Task:
|
||||
"""任务数据类"""
|
||||
task_id: str
|
||||
task_type: str
|
||||
status: TaskStatus
|
||||
created_at: datetime
|
||||
updated_at: datetime
|
||||
progress: int = 0 # 总进度百分比 0-100
|
||||
message: str = "" # 状态消息
|
||||
result: Optional[Dict] = None # 任务结果
|
||||
error: Optional[str] = None # 错误信息
|
||||
metadata: Dict = field(default_factory=dict) # 额外元数据
|
||||
progress_detail: Dict = field(default_factory=dict) # 详细进度信息
|
||||
|
||||
progress: int = 0
|
||||
message: str = ""
|
||||
result: Optional[Dict] = None
|
||||
error: Optional[str] = None
|
||||
metadata: Dict = field(default_factory=dict)
|
||||
progress_detail: Dict = field(default_factory=dict)
|
||||
|
||||
def to_dict(self) -> Dict[str, Any]:
|
||||
"""转换为字典"""
|
||||
return {
|
||||
"task_id": self.task_id,
|
||||
"task_type": self.task_type,
|
||||
@@ -54,57 +59,182 @@ class Task:
|
||||
|
||||
|
||||
class TaskManager:
|
||||
"""
|
||||
任务管理器
|
||||
线程安全的任务状态管理
|
||||
"""
|
||||
|
||||
_instance = None
|
||||
_lock = threading.Lock()
|
||||
|
||||
def __new__(cls):
|
||||
"""单例模式"""
|
||||
if cls._instance is None:
|
||||
with cls._lock:
|
||||
if cls._instance is None:
|
||||
cls._instance = super().__new__(cls)
|
||||
cls._instance._tasks: Dict[str, Task] = {}
|
||||
cls._instance._task_lock = threading.Lock()
|
||||
return cls._instance
|
||||
|
||||
"""Thread-safe task facade bound to one app/session factory."""
|
||||
|
||||
_configured_session_factory = None
|
||||
_config_lock = threading.Lock()
|
||||
_fallback_tasks: Dict[str, Task] = {}
|
||||
_fallback_lock = threading.Lock()
|
||||
|
||||
def __init__(self, session_factory=None):
|
||||
"""Bind this manager to an explicit or current-app session factory.
|
||||
|
||||
A manager created during a request/app context keeps that app's factory
|
||||
for background work, but it cannot be reused inside a different Flask
|
||||
app. The explicit class configuration remains only for legacy tests and
|
||||
callers that run outside Flask.
|
||||
"""
|
||||
self._bound_app = None
|
||||
if has_app_context():
|
||||
app = cast(Any, current_app)._get_current_object()
|
||||
current_factory = app.extensions.get("crowdsight_session_factory")
|
||||
if not callable(current_factory):
|
||||
raise RuntimeError("task_session_factory_required")
|
||||
if session_factory is not None and session_factory is not current_factory:
|
||||
raise RuntimeError("task_session_factory_mismatch")
|
||||
bound_factory = current_factory
|
||||
self._bound_app = app
|
||||
elif session_factory is not None:
|
||||
bound_factory = session_factory
|
||||
else:
|
||||
# No Flask app context: use the explicit legacy/test binding.
|
||||
bound_factory = type(self)._configured_session_factory
|
||||
|
||||
self._session_factory = bound_factory
|
||||
self._tasks = type(self)._fallback_tasks
|
||||
self._task_lock = type(self)._fallback_lock
|
||||
|
||||
@classmethod
|
||||
def configure(cls, session_factory) -> None:
|
||||
"""Set an explicit outside-Flask binding for tests/legacy adapters."""
|
||||
with cls._config_lock:
|
||||
cls._configured_session_factory = session_factory
|
||||
with cls._fallback_lock:
|
||||
cls._fallback_tasks.clear()
|
||||
|
||||
def _factory(self):
|
||||
if not has_app_context():
|
||||
return self._session_factory
|
||||
|
||||
app = cast(Any, current_app)._get_current_object()
|
||||
current_factory = app.extensions.get("crowdsight_session_factory")
|
||||
if not callable(current_factory):
|
||||
raise RuntimeError("task_session_factory_required")
|
||||
if self._bound_app is not None and self._bound_app is not app:
|
||||
raise RuntimeError("task_app_context_mismatch")
|
||||
if self._session_factory is not None and self._session_factory is not current_factory:
|
||||
raise RuntimeError("task_session_factory_mismatch")
|
||||
if self._session_factory is None:
|
||||
self._session_factory = current_factory
|
||||
self._bound_app = app
|
||||
return self._session_factory
|
||||
|
||||
@staticmethod
|
||||
def _bounded_text(value: Any, limit: int = 4000) -> str:
|
||||
if value is None:
|
||||
return ""
|
||||
return str(value)[:limit]
|
||||
|
||||
@staticmethod
|
||||
def _job_status(status: TaskStatus | str | None) -> str | None:
|
||||
if status is None:
|
||||
return None
|
||||
value = status.value if isinstance(status, TaskStatus) else str(status)
|
||||
return {
|
||||
TaskStatus.PENDING.value: JobStatus.QUEUED.value,
|
||||
TaskStatus.PROCESSING.value: JobStatus.RUNNING.value,
|
||||
TaskStatus.COMPLETED.value: JobStatus.SUCCEEDED.value,
|
||||
TaskStatus.FAILED.value: JobStatus.FAILED.value,
|
||||
JobStatus.QUEUED.value: JobStatus.QUEUED.value,
|
||||
JobStatus.RUNNING.value: JobStatus.RUNNING.value,
|
||||
JobStatus.SUCCEEDED.value: JobStatus.SUCCEEDED.value,
|
||||
JobStatus.FAILED.value: JobStatus.FAILED.value,
|
||||
JobStatus.CANCELLED.value: JobStatus.CANCELLED.value,
|
||||
}.get(value)
|
||||
|
||||
@staticmethod
|
||||
def _task_status(status: str | JobStatus) -> TaskStatus:
|
||||
value = status.value if isinstance(status, JobStatus) else str(status)
|
||||
return {
|
||||
JobStatus.QUEUED.value: TaskStatus.PENDING,
|
||||
JobStatus.RUNNING.value: TaskStatus.PROCESSING,
|
||||
JobStatus.SUCCEEDED.value: TaskStatus.COMPLETED,
|
||||
JobStatus.FAILED.value: TaskStatus.FAILED,
|
||||
JobStatus.CANCELLED.value: TaskStatus.FAILED,
|
||||
}.get(value, TaskStatus.FAILED)
|
||||
|
||||
@classmethod
|
||||
def _from_job(cls, job: Job) -> Task:
|
||||
metadata = job.job_metadata if isinstance(job.job_metadata, dict) else {}
|
||||
result = job.result if isinstance(job.result, dict) else job.result
|
||||
detail = job.progress_detail if isinstance(job.progress_detail, dict) else {}
|
||||
return Task(
|
||||
task_id=job.id,
|
||||
task_type=job.operation,
|
||||
status=cls._task_status(job.status),
|
||||
created_at=job.created_at,
|
||||
updated_at=job.updated_at,
|
||||
progress=job.progress,
|
||||
message=job.message,
|
||||
result=result,
|
||||
error=job.error_code,
|
||||
metadata=metadata,
|
||||
progress_detail=detail,
|
||||
)
|
||||
|
||||
def create_task(self, task_type: str, metadata: Optional[Dict] = None) -> str:
|
||||
"""
|
||||
创建新任务
|
||||
|
||||
Args:
|
||||
task_type: 任务类型
|
||||
metadata: 额外元数据
|
||||
|
||||
Returns:
|
||||
任务ID
|
||||
"""
|
||||
metadata = metadata or {}
|
||||
factory = self._factory()
|
||||
if factory is not None:
|
||||
organization_id = metadata.get("organization_id")
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
raise ValueError("task_scope_required")
|
||||
with factory() as session:
|
||||
job = Job(
|
||||
organization_id=organization_id,
|
||||
owner_user_id=metadata.get("owner_user_id"),
|
||||
project_id=metadata.get("project_id"),
|
||||
graph_id=metadata.get("graph_id"),
|
||||
operation=self._bounded_text(task_type, 120),
|
||||
status=JobStatus.QUEUED.value,
|
||||
job_metadata=metadata,
|
||||
progress_detail={},
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
return job.id
|
||||
|
||||
task_id = str(uuid.uuid4())
|
||||
now = datetime.now()
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
task = Task(
|
||||
task_id=task_id,
|
||||
task_type=task_type,
|
||||
status=TaskStatus.PENDING,
|
||||
created_at=now,
|
||||
updated_at=now,
|
||||
metadata=metadata or {}
|
||||
metadata=metadata,
|
||||
)
|
||||
|
||||
with self._task_lock:
|
||||
self._tasks[task_id] = task
|
||||
|
||||
return task_id
|
||||
|
||||
def get_task(self, task_id: str) -> Optional[Task]:
|
||||
"""获取任务"""
|
||||
|
||||
def get_task(
|
||||
self,
|
||||
task_id: str,
|
||||
*,
|
||||
organization_id: Optional[str] = None,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> Optional[Task]:
|
||||
factory = self._factory()
|
||||
if factory is not None:
|
||||
with factory() as session:
|
||||
statement = select(Job).where(Job.id == task_id)
|
||||
if organization_id is not None:
|
||||
statement = statement.where(Job.organization_id == organization_id)
|
||||
if owner_user_id is not None:
|
||||
statement = statement.where(Job.owner_user_id == owner_user_id)
|
||||
job = session.scalar(statement)
|
||||
return self._from_job(job) if job is not None else None
|
||||
with self._task_lock:
|
||||
return self._tasks.get(task_id)
|
||||
|
||||
task = self._tasks.get(task_id)
|
||||
if task is None:
|
||||
return None
|
||||
if organization_id is not None and task.metadata.get("organization_id") != organization_id:
|
||||
return None
|
||||
if owner_user_id is not None and task.metadata.get("owner_user_id") != owner_user_id:
|
||||
return None
|
||||
return task
|
||||
|
||||
def update_task(
|
||||
self,
|
||||
task_id: str,
|
||||
@@ -113,28 +243,45 @@ class TaskManager:
|
||||
message: Optional[str] = None,
|
||||
result: Optional[Dict] = None,
|
||||
error: Optional[str] = None,
|
||||
progress_detail: Optional[Dict] = None
|
||||
progress_detail: Optional[Dict] = None,
|
||||
):
|
||||
"""
|
||||
更新任务状态
|
||||
|
||||
Args:
|
||||
task_id: 任务ID
|
||||
status: 新状态
|
||||
progress: 进度
|
||||
message: 消息
|
||||
result: 结果
|
||||
error: 错误信息
|
||||
progress_detail: 详细进度信息
|
||||
"""
|
||||
factory = self._factory()
|
||||
if factory is not None:
|
||||
with factory() as session:
|
||||
job = session.get(Job, task_id)
|
||||
if job is None:
|
||||
return
|
||||
mapped_status = self._job_status(status)
|
||||
if mapped_status is not None:
|
||||
job.status = mapped_status
|
||||
if mapped_status in {
|
||||
JobStatus.SUCCEEDED.value,
|
||||
JobStatus.FAILED.value,
|
||||
JobStatus.CANCELLED.value,
|
||||
}:
|
||||
job.finished_at = datetime.now(timezone.utc)
|
||||
if progress is not None:
|
||||
job.progress = min(max(int(progress), 0), 100)
|
||||
if message is not None:
|
||||
job.message = self._bounded_text(message)
|
||||
if result is not None:
|
||||
job.result = result
|
||||
if error is not None:
|
||||
job.error_code = self._bounded_text(error, 120)
|
||||
if progress_detail is not None:
|
||||
job.progress_detail = progress_detail
|
||||
job.updated_at = datetime.now(timezone.utc)
|
||||
session.commit()
|
||||
return
|
||||
|
||||
with self._task_lock:
|
||||
task = self._tasks.get(task_id)
|
||||
if task:
|
||||
task.updated_at = datetime.now()
|
||||
task.updated_at = datetime.now(timezone.utc)
|
||||
if status is not None:
|
||||
task.status = status
|
||||
if progress is not None:
|
||||
task.progress = progress
|
||||
task.progress = min(max(int(progress), 0), 100)
|
||||
if message is not None:
|
||||
task.message = message
|
||||
if result is not None:
|
||||
@@ -143,44 +290,71 @@ class TaskManager:
|
||||
task.error = error
|
||||
if progress_detail is not None:
|
||||
task.progress_detail = progress_detail
|
||||
|
||||
|
||||
def complete_task(self, task_id: str, result: Dict):
|
||||
"""标记任务完成"""
|
||||
self.update_task(
|
||||
task_id,
|
||||
status=TaskStatus.COMPLETED,
|
||||
progress=100,
|
||||
message=t('progress.taskComplete'),
|
||||
result=result
|
||||
message=t("progress.taskComplete"),
|
||||
result=result,
|
||||
)
|
||||
|
||||
|
||||
def fail_task(self, task_id: str, error: str):
|
||||
"""标记任务失败"""
|
||||
self.update_task(
|
||||
task_id,
|
||||
status=TaskStatus.FAILED,
|
||||
message=t('progress.taskFailed'),
|
||||
error=error
|
||||
message=t("progress.taskFailed"),
|
||||
error=error,
|
||||
)
|
||||
|
||||
def list_tasks(self, task_type: Optional[str] = None) -> list:
|
||||
"""列出任务"""
|
||||
|
||||
def list_tasks(
|
||||
self,
|
||||
task_type: Optional[str] = None,
|
||||
*,
|
||||
organization_id: Optional[str] = None,
|
||||
owner_user_id: Optional[str] = None,
|
||||
) -> list:
|
||||
factory = self._factory()
|
||||
if factory is not None:
|
||||
with factory() as session:
|
||||
statement = select(Job).order_by(Job.created_at.desc())
|
||||
if task_type:
|
||||
statement = statement.where(Job.operation == task_type)
|
||||
if organization_id is not None:
|
||||
statement = statement.where(Job.organization_id == organization_id)
|
||||
if owner_user_id is not None:
|
||||
statement = statement.where(Job.owner_user_id == owner_user_id)
|
||||
jobs = session.scalars(statement.limit(100)).all()
|
||||
return [self._from_job(job) for job in jobs]
|
||||
with self._task_lock:
|
||||
tasks = list(self._tasks.values())
|
||||
if task_type:
|
||||
tasks = [t for t in tasks if t.task_type == task_type]
|
||||
return [t.to_dict() for t in sorted(tasks, key=lambda x: x.created_at, reverse=True)]
|
||||
|
||||
tasks = [task for task in tasks if task.task_type == task_type]
|
||||
if organization_id is not None:
|
||||
tasks = [task for task in tasks if task.metadata.get("organization_id") == organization_id]
|
||||
if owner_user_id is not None:
|
||||
tasks = [task for task in tasks if task.metadata.get("owner_user_id") == owner_user_id]
|
||||
return [task for task in sorted(tasks, key=lambda item: item.created_at, reverse=True)]
|
||||
|
||||
def cleanup_old_tasks(self, max_age_hours: int = 24):
|
||||
"""清理旧任务"""
|
||||
from datetime import timedelta
|
||||
cutoff = datetime.now() - timedelta(hours=max_age_hours)
|
||||
|
||||
cutoff = datetime.now(timezone.utc) - timedelta(hours=max_age_hours)
|
||||
factory = self._factory()
|
||||
if factory is not None:
|
||||
with factory() as session:
|
||||
session.execute(
|
||||
delete(Job).where(
|
||||
Job.created_at < cutoff,
|
||||
Job.status.in_([JobStatus.SUCCEEDED.value, JobStatus.FAILED.value]),
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
return
|
||||
with self._task_lock:
|
||||
old_ids = [
|
||||
tid for tid, task in self._tasks.items()
|
||||
task_id
|
||||
for task_id, task in self._tasks.items()
|
||||
if task.created_at < cutoff and task.status in [TaskStatus.COMPLETED, TaskStatus.FAILED]
|
||||
]
|
||||
for tid in old_ids:
|
||||
del self._tasks[tid]
|
||||
|
||||
for task_id in old_ids:
|
||||
del self._tasks[task_id]
|
||||
|
||||
45
backend/app/models/usage.py
Normal file
45
backend/app/models/usage.py
Normal file
@@ -0,0 +1,45 @@
|
||||
"""Durable LLM usage/cost events."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from uuid import uuid4
|
||||
|
||||
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, func, text
|
||||
from sqlalchemy.orm import Mapped, mapped_column
|
||||
|
||||
from ..db import Base
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class UsageEvent(Base):
|
||||
"""One LLM usage record. Never stores prompt content or secrets."""
|
||||
|
||||
__tablename__ = "usage_events"
|
||||
__table_args__ = (
|
||||
Index("ix_usage_org_created", "organization_id", "created_at"),
|
||||
Index("ix_usage_org_user", "organization_id", "user_id"),
|
||||
)
|
||||
|
||||
id: Mapped[str] = mapped_column(
|
||||
String(64), primary_key=True, default=lambda: f"usage_{uuid4().hex}"
|
||||
)
|
||||
organization_id: Mapped[str] = mapped_column(
|
||||
ForeignKey("organizations.id", ondelete="CASCADE"), nullable=False, index=True
|
||||
)
|
||||
user_id: Mapped[str | None] = mapped_column(
|
||||
ForeignKey("users.id", ondelete="SET NULL"), nullable=True, index=True
|
||||
)
|
||||
operation: Mapped[str] = mapped_column(String(160), nullable=False)
|
||||
model: Mapped[str | None] = mapped_column(String(120), nullable=True)
|
||||
input_tokens: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default=text("0"))
|
||||
output_tokens: Mapped[int] = mapped_column(Integer, nullable=False, default=0, server_default=text("0"))
|
||||
estimated_cost: Mapped[float] = mapped_column(
|
||||
Float, nullable=False, default=0.0, server_default=text("0")
|
||||
)
|
||||
created_at: Mapped[datetime] = mapped_column(
|
||||
DateTime(timezone=True), nullable=False, default=_utc_now, server_default=func.now()
|
||||
)
|
||||
19
backend/app/security/__init__.py
Normal file
19
backend/app/security/__init__.py
Normal file
@@ -0,0 +1,19 @@
|
||||
"""Security policy package."""
|
||||
|
||||
from .policy import (
|
||||
Actor,
|
||||
AuthorizationError,
|
||||
Role,
|
||||
assert_can_access_resource,
|
||||
assert_can_manage_llm_settings,
|
||||
assert_can_manage_user,
|
||||
)
|
||||
|
||||
__all__ = [
|
||||
"Actor",
|
||||
"AuthorizationError",
|
||||
"Role",
|
||||
"assert_can_access_resource",
|
||||
"assert_can_manage_llm_settings",
|
||||
"assert_can_manage_user",
|
||||
]
|
||||
108
backend/app/security/auth.py
Normal file
108
backend/app/security/auth.py
Normal file
@@ -0,0 +1,108 @@
|
||||
"""Flask request authentication decorators for tenant-scoped routes."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hmac
|
||||
import secrets
|
||||
from functools import wraps
|
||||
|
||||
from flask import current_app, g, request
|
||||
from itsdangerous import BadSignature, URLSafeTimedSerializer
|
||||
|
||||
from ..services.identity import SessionService
|
||||
from ..utils.api_errors import ApiError
|
||||
from .policy import Actor, Role
|
||||
|
||||
|
||||
def _session_factory():
|
||||
factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if factory is None:
|
||||
raise ApiError("auth_unavailable", 503, "api.internalError")
|
||||
return factory
|
||||
|
||||
|
||||
def _csrf_serializer() -> URLSafeTimedSerializer:
|
||||
secret_key = current_app.secret_key
|
||||
if not secret_key:
|
||||
raise ApiError("auth_unavailable", 503, "api.internalError")
|
||||
return URLSafeTimedSerializer(secret_key, salt="crowdsight-csrf")
|
||||
|
||||
|
||||
def issue_csrf_token() -> str:
|
||||
return _csrf_serializer().dumps(secrets.token_urlsafe(24))
|
||||
|
||||
|
||||
def _validate_csrf() -> None:
|
||||
if request.method in {"GET", "HEAD", "OPTIONS"}:
|
||||
return
|
||||
cookie_token = request.cookies.get("crowdsight_csrf", "")
|
||||
header_token = request.headers.get("X-CSRF-Token", "")
|
||||
if not cookie_token or not header_token or not hmac.compare_digest(cookie_token, header_token):
|
||||
raise ApiError("csrf_failed", 403, "common.error")
|
||||
try:
|
||||
_csrf_serializer().loads(cookie_token, max_age=SessionService.DEFAULT_TTL_SECONDS)
|
||||
except BadSignature as exc:
|
||||
raise ApiError("csrf_failed", 403, "common.error") from exc
|
||||
|
||||
|
||||
def authenticate_readonly_request() -> None:
|
||||
"""Authenticate a legacy blueprint without exposing a DB session to the route."""
|
||||
if getattr(g, "auth_context", None) is not None:
|
||||
return
|
||||
raw_token = request.cookies.get("crowdsight_session", "")
|
||||
with _session_factory()() as db_session:
|
||||
context = SessionService.resolve(db_session, raw_token)
|
||||
if context is None:
|
||||
raise ApiError("unauthorized", 401, "common.unauthorized")
|
||||
_validate_csrf()
|
||||
g.auth_context = context
|
||||
|
||||
|
||||
def current_actor() -> Actor:
|
||||
context = getattr(g, "auth_context", None)
|
||||
if context is None:
|
||||
raise ApiError("unauthorized", 401, "common.unauthorized")
|
||||
return Actor(
|
||||
user_id=context.user.id,
|
||||
organization_id=context.organization.id,
|
||||
role=context.membership.role,
|
||||
)
|
||||
|
||||
|
||||
def require_auth(view):
|
||||
@wraps(view)
|
||||
def wrapped(*args, **kwargs):
|
||||
raw_token = request.cookies.get("crowdsight_session", "")
|
||||
with _session_factory()() as db_session:
|
||||
context = SessionService.resolve(db_session, raw_token)
|
||||
if context is None:
|
||||
raise ApiError("unauthorized", 401, "common.unauthorized")
|
||||
_validate_csrf()
|
||||
|
||||
g.auth_context = context
|
||||
g.db_session = db_session
|
||||
try:
|
||||
response = view(*args, **kwargs)
|
||||
db_session.commit()
|
||||
return response
|
||||
except Exception:
|
||||
db_session.rollback()
|
||||
raise
|
||||
|
||||
return wrapped
|
||||
|
||||
|
||||
def require_roles(*allowed_roles: Role | str):
|
||||
allowed = {role if isinstance(role, Role) else Role(role) for role in allowed_roles}
|
||||
|
||||
def decorator(view):
|
||||
@wraps(view)
|
||||
def wrapped(*args, **kwargs):
|
||||
actor = current_actor()
|
||||
if actor.role not in allowed:
|
||||
raise ApiError("forbidden", 403, "common.error")
|
||||
return view(*args, **kwargs)
|
||||
|
||||
return require_auth(wrapped)
|
||||
|
||||
return decorator
|
||||
121
backend/app/security/policy.py
Normal file
121
backend/app/security/policy.py
Normal file
@@ -0,0 +1,121 @@
|
||||
"""Fail-closed authorization primitives for tenant-scoped resource services.
|
||||
|
||||
This module deliberately has no Flask or database dependency. Route handlers and
|
||||
repositories can use the same policy contract without importing the whole app.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from enum import Enum
|
||||
from typing import Any
|
||||
|
||||
|
||||
class Role(str, Enum):
|
||||
SUPER_ADMIN = "super_admin"
|
||||
ADMIN = "admin"
|
||||
USER = "user"
|
||||
|
||||
|
||||
class AuthorizationError(PermissionError):
|
||||
"""Raised when an actor cannot perform a requested operation."""
|
||||
|
||||
def __init__(self, code: str = "forbidden"):
|
||||
self.code = code
|
||||
super().__init__(code)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Actor:
|
||||
"""Minimal authenticated identity required by policy checks."""
|
||||
|
||||
user_id: str
|
||||
organization_id: str
|
||||
role: Role | str
|
||||
|
||||
|
||||
def _role(actor: Actor) -> Role:
|
||||
try:
|
||||
return actor.role if isinstance(actor.role, Role) else Role(actor.role)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AuthorizationError("invalid_actor_role") from exc
|
||||
|
||||
|
||||
def _require_non_empty(value: Any, code: str) -> str:
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise AuthorizationError(code)
|
||||
return value
|
||||
|
||||
|
||||
def assert_can_access_resource(
|
||||
actor: Actor,
|
||||
*,
|
||||
resource_organization_id: str,
|
||||
owner_user_id: str | None,
|
||||
action: str = "read",
|
||||
platform_scope: bool = False,
|
||||
) -> None:
|
||||
"""Raise unless ``actor`` may access a tenant-owned resource.
|
||||
|
||||
``platform_scope`` is explicit even for super admins so a caller cannot
|
||||
accidentally turn every ordinary resource lookup into a cross-tenant path.
|
||||
A regular user must have an exact owner marker; missing/falsey ownership is
|
||||
denied rather than treated as public.
|
||||
"""
|
||||
del action # The first contract is scope; action-specific rules layer on it.
|
||||
role = _role(actor)
|
||||
resource_org = _require_non_empty(resource_organization_id, "invalid_resource_scope")
|
||||
actor_org = _require_non_empty(actor.organization_id, "invalid_actor_scope")
|
||||
|
||||
if role is Role.SUPER_ADMIN:
|
||||
if platform_scope or actor_org == resource_org:
|
||||
return
|
||||
raise AuthorizationError("cross_tenant_scope_required")
|
||||
|
||||
if actor_org != resource_org:
|
||||
raise AuthorizationError("cross_tenant_forbidden")
|
||||
|
||||
if role is Role.ADMIN:
|
||||
return
|
||||
|
||||
if role is Role.USER and owner_user_id == actor.user_id and actor.user_id:
|
||||
return
|
||||
|
||||
raise AuthorizationError("resource_owner_required")
|
||||
|
||||
|
||||
def assert_can_manage_user(
|
||||
actor: Actor,
|
||||
*,
|
||||
target_organization_id: str,
|
||||
target_role: Role | str,
|
||||
platform_scope: bool = False,
|
||||
) -> None:
|
||||
"""Raise unless ``actor`` may manage a target account/membership."""
|
||||
role = _role(actor)
|
||||
target_org = _require_non_empty(target_organization_id, "invalid_target_scope")
|
||||
try:
|
||||
requested_role = target_role if isinstance(target_role, Role) else Role(target_role)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise AuthorizationError("invalid_target_role") from exc
|
||||
|
||||
if role is Role.SUPER_ADMIN:
|
||||
if platform_scope or actor.organization_id == target_org:
|
||||
return
|
||||
raise AuthorizationError("cross_tenant_scope_required")
|
||||
|
||||
if role is Role.ADMIN:
|
||||
if actor.organization_id != target_org:
|
||||
raise AuthorizationError("cross_tenant_forbidden")
|
||||
if requested_role is Role.USER:
|
||||
return
|
||||
raise AuthorizationError("admin_role_grant_forbidden")
|
||||
|
||||
raise AuthorizationError("user_management_forbidden")
|
||||
|
||||
|
||||
def assert_can_manage_llm_settings(actor: Actor, *, platform_scope: bool = False) -> None:
|
||||
"""Only a super admin with explicit platform scope may mutate LLM settings."""
|
||||
if _role(actor) is Role.SUPER_ADMIN and platform_scope:
|
||||
return
|
||||
raise AuthorizationError("llm_settings_forbidden")
|
||||
223
backend/app/security/resources.py
Normal file
223
backend/app/security/resources.py
Normal file
@@ -0,0 +1,223 @@
|
||||
"""Fail-closed tenant/owner lookup helpers for file-backed resources with
|
||||
durable (PostgreSQL-target) read-first cutover for projects.
|
||||
|
||||
Project lookups prefer the durable ``projects`` table when the session factory
|
||||
is available, then fall back to the legacy filesystem ``ProjectManager``. This
|
||||
keeps read paths working during the filesystem → durable migration without
|
||||
breaking existing routes that consume the legacy ``Project`` dataclass shape.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
from flask import current_app, request
|
||||
|
||||
from ..models.project import Project, ProjectManager, ProjectStatus
|
||||
from ..models.task import Task, TaskManager
|
||||
from ..services.report_agent import Report, ReportManager
|
||||
from ..services.simulation_manager import SimulationManager, SimulationState
|
||||
from ..utils.api_errors import ApiError
|
||||
from .auth import current_actor
|
||||
from .policy import Role
|
||||
|
||||
_SAFE_RESOURCE_ID = re.compile(r"^[A-Za-z0-9][A-Za-z0-9_-]{0,127}$")
|
||||
|
||||
|
||||
def _valid_resource_id(value: str) -> bool:
|
||||
return isinstance(value, str) and bool(_SAFE_RESOURCE_ID.fullmatch(value))
|
||||
|
||||
|
||||
def _durable_project_for_scope(
|
||||
project_id: str, *, organization_id: str, owner_user_id: Optional[str]
|
||||
) -> Optional[Project]:
|
||||
"""Resolve a project from the durable table and wrap it in the legacy shape.
|
||||
|
||||
Only active when the local (durable) memory backend is selected; the legacy
|
||||
Zep backend keeps reading the filesystem manager exclusively. Returns None
|
||||
when the durable store is unavailable or no matching tenant/owner row
|
||||
exists. Does not raise.
|
||||
"""
|
||||
try:
|
||||
from ..config import Config
|
||||
|
||||
if Config.MEMORY_BACKEND != "local":
|
||||
return None
|
||||
session_factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if session_factory is None:
|
||||
return None
|
||||
from ..services.product_repository import ProductRepository
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
row = ProductRepository(session).get_project(
|
||||
project_id, organization_id=organization_id
|
||||
)
|
||||
finally:
|
||||
session.close()
|
||||
if row is None:
|
||||
return None
|
||||
if owner_user_id is not None and row.owner_user_id != owner_user_id:
|
||||
return None
|
||||
try:
|
||||
project_status = ProjectStatus(row.status)
|
||||
except ValueError:
|
||||
project_status = ProjectStatus.CREATED
|
||||
return Project(
|
||||
project_id=row.id,
|
||||
name=row.name,
|
||||
status=project_status,
|
||||
created_at=row.created_at.isoformat() if row.created_at else "",
|
||||
updated_at=row.updated_at.isoformat() if row.updated_at else "",
|
||||
organization_id=row.organization_id,
|
||||
owner_user_id=row.owner_user_id,
|
||||
files=[],
|
||||
total_text_length=row.total_text_length,
|
||||
ontology=row.ontology or None,
|
||||
analysis_summary=row.analysis_summary,
|
||||
graph_id=row.graph_id,
|
||||
graph_build_task_id=row.graph_build_task_id,
|
||||
simulation_requirement=row.simulation_requirement,
|
||||
error=row.error,
|
||||
)
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
|
||||
def scoped_project(project_id: str) -> Optional[Project]:
|
||||
if not _valid_resource_id(project_id):
|
||||
return None
|
||||
actor = current_actor()
|
||||
owner_user_id = actor.user_id if actor.role is Role.USER else None
|
||||
durable = _durable_project_for_scope(
|
||||
project_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
if durable is not None:
|
||||
return durable
|
||||
return ProjectManager.get_project_for_scope(
|
||||
project_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
)
|
||||
|
||||
|
||||
def require_scoped_project(project_id: str) -> Project:
|
||||
project = scoped_project(project_id)
|
||||
if project is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
return project
|
||||
|
||||
|
||||
def scoped_graph(graph_id: str) -> Optional[Project]:
|
||||
if not _valid_resource_id(graph_id):
|
||||
return None
|
||||
actor = current_actor()
|
||||
return ProjectManager.find_project_by_graph_id(
|
||||
graph_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id if actor.role is Role.USER else None,
|
||||
)
|
||||
|
||||
|
||||
def scoped_simulation(simulation_id: str) -> Optional[SimulationState]:
|
||||
if not _valid_resource_id(simulation_id):
|
||||
return None
|
||||
state = SimulationManager().get_simulation(simulation_id)
|
||||
if state is None:
|
||||
return None
|
||||
if scoped_project(state.project_id) is None:
|
||||
return None
|
||||
return state
|
||||
|
||||
|
||||
def require_scoped_simulation(simulation_id: str) -> SimulationState:
|
||||
state = scoped_simulation(simulation_id)
|
||||
if state is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
return state
|
||||
|
||||
|
||||
def scoped_task(task_id: str) -> Optional[Task]:
|
||||
if not _valid_resource_id(task_id):
|
||||
return None
|
||||
actor = current_actor()
|
||||
task = TaskManager().get_task(
|
||||
task_id,
|
||||
organization_id=actor.organization_id,
|
||||
owner_user_id=actor.user_id if actor.role is Role.USER else None,
|
||||
)
|
||||
if task is None:
|
||||
return None
|
||||
metadata = task.metadata if isinstance(task.metadata, dict) else {}
|
||||
if metadata.get("organization_id") != actor.organization_id:
|
||||
return None
|
||||
if actor.role is Role.USER and metadata.get("owner_user_id") != actor.user_id:
|
||||
return None
|
||||
return task
|
||||
|
||||
|
||||
def require_scoped_task(task_id: str) -> Task:
|
||||
task = scoped_task(task_id)
|
||||
if task is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
return task
|
||||
|
||||
|
||||
def scoped_report(report_id: str) -> Optional[Report]:
|
||||
if not _valid_resource_id(report_id):
|
||||
return None
|
||||
report = ReportManager.get_report(report_id)
|
||||
if report is None:
|
||||
return None
|
||||
if scoped_simulation(report.simulation_id) is None:
|
||||
return None
|
||||
return report
|
||||
|
||||
|
||||
def require_scoped_report(report_id: str) -> Report:
|
||||
report = scoped_report(report_id)
|
||||
if report is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
return report
|
||||
|
||||
|
||||
def scoped_reports(simulation_id: str | None = None, limit: int = 50) -> list[Report]:
|
||||
if simulation_id is not None and scoped_simulation(simulation_id) is None:
|
||||
return []
|
||||
safe_limit = min(max(int(limit), 1), 100)
|
||||
reports = ReportManager.list_reports(simulation_id=simulation_id, limit=safe_limit * 4)
|
||||
return [report for report in reports if scoped_report(report.report_id) is not None][:safe_limit]
|
||||
|
||||
|
||||
def scoped_simulations(project_id: str | None = None) -> list[SimulationState]:
|
||||
if project_id is not None and scoped_project(project_id) is None:
|
||||
return []
|
||||
simulations = SimulationManager().list_simulations(project_id=project_id)
|
||||
return [state for state in simulations if scoped_simulation(state.simulation_id) is not None]
|
||||
|
||||
|
||||
def enforce_request_scope() -> None:
|
||||
"""Reject IDs outside the current actor's tenant/owner scope before handlers run."""
|
||||
view_args = request.view_args or {}
|
||||
payload = request.get_json(silent=True)
|
||||
payload = payload if isinstance(payload, dict) else {}
|
||||
|
||||
project_id = view_args.get("project_id") or request.args.get("project_id") or payload.get("project_id")
|
||||
simulation_id = view_args.get("simulation_id") or request.args.get("simulation_id") or payload.get("simulation_id")
|
||||
report_id = view_args.get("report_id") or request.args.get("report_id") or payload.get("report_id")
|
||||
graph_id = view_args.get("graph_id") or request.args.get("graph_id") or payload.get("graph_id")
|
||||
task_id = view_args.get("task_id") or request.args.get("task_id") or payload.get("task_id")
|
||||
|
||||
if project_id and scoped_project(project_id) is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
if simulation_id and scoped_simulation(simulation_id) is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
if report_id and scoped_report(report_id) is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
if graph_id and scoped_graph(graph_id) is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
if task_id and scoped_task(task_id) is None:
|
||||
raise ApiError("resource_not_found", 404, "common.notFound")
|
||||
@@ -1,73 +1,61 @@
|
||||
"""
|
||||
业务服务模块
|
||||
"""Business service exports with lazy imports for backend isolation.
|
||||
|
||||
The local memory backend must be importable without eagerly loading Zep-only
|
||||
consumers. Legacy package-level names remain available through ``__getattr__``
|
||||
when a caller explicitly requests them.
|
||||
"""
|
||||
|
||||
from .ontology_generator import OntologyGenerator
|
||||
from .graph_builder import GraphBuilderService
|
||||
from .text_processor import TextProcessor
|
||||
from .zep_entity_reader import ZepEntityReader, EntityNode, FilteredEntities
|
||||
from .oasis_profile_generator import OasisProfileGenerator, OasisAgentProfile
|
||||
from .simulation_manager import SimulationManager, SimulationState, SimulationStatus
|
||||
from .simulation_config_generator import (
|
||||
SimulationConfigGenerator,
|
||||
SimulationParameters,
|
||||
AgentActivityConfig,
|
||||
TimeSimulationConfig,
|
||||
EventConfig,
|
||||
PlatformConfig
|
||||
)
|
||||
from .simulation_runner import (
|
||||
SimulationRunner,
|
||||
SimulationRunState,
|
||||
RunnerStatus,
|
||||
AgentAction,
|
||||
RoundSummary
|
||||
)
|
||||
from .zep_graph_memory_updater import (
|
||||
ZepGraphMemoryUpdater,
|
||||
ZepGraphMemoryManager,
|
||||
AgentActivity
|
||||
)
|
||||
from .simulation_ipc import (
|
||||
SimulationIPCClient,
|
||||
SimulationIPCServer,
|
||||
IPCCommand,
|
||||
IPCResponse,
|
||||
CommandType,
|
||||
CommandStatus
|
||||
)
|
||||
from importlib import import_module
|
||||
|
||||
__all__ = [
|
||||
'OntologyGenerator',
|
||||
'GraphBuilderService',
|
||||
'TextProcessor',
|
||||
'ZepEntityReader',
|
||||
'EntityNode',
|
||||
'FilteredEntities',
|
||||
'OasisProfileGenerator',
|
||||
'OasisAgentProfile',
|
||||
'SimulationManager',
|
||||
'SimulationState',
|
||||
'SimulationStatus',
|
||||
'SimulationConfigGenerator',
|
||||
'SimulationParameters',
|
||||
'AgentActivityConfig',
|
||||
'TimeSimulationConfig',
|
||||
'EventConfig',
|
||||
'PlatformConfig',
|
||||
'SimulationRunner',
|
||||
'SimulationRunState',
|
||||
'RunnerStatus',
|
||||
'AgentAction',
|
||||
'RoundSummary',
|
||||
'ZepGraphMemoryUpdater',
|
||||
'ZepGraphMemoryManager',
|
||||
'AgentActivity',
|
||||
'SimulationIPCClient',
|
||||
'SimulationIPCServer',
|
||||
'IPCCommand',
|
||||
'IPCResponse',
|
||||
'CommandType',
|
||||
'CommandStatus',
|
||||
]
|
||||
_EXPORTS = {
|
||||
"OntologyGenerator": (".ontology_generator", "OntologyGenerator"),
|
||||
"GraphBuilderService": (".graph_builder", "GraphBuilderService"),
|
||||
"TextProcessor": (".text_processor", "TextProcessor"),
|
||||
"ZepEntityReader": (".zep_entity_reader", "ZepEntityReader"),
|
||||
"EntityNode": (".zep_entity_reader", "EntityNode"),
|
||||
"FilteredEntities": (".zep_entity_reader", "FilteredEntities"),
|
||||
"OasisProfileGenerator": (".oasis_profile_generator", "OasisProfileGenerator"),
|
||||
"OasisAgentProfile": (".oasis_profile_generator", "OasisAgentProfile"),
|
||||
"SimulationManager": (".simulation_manager", "SimulationManager"),
|
||||
"SimulationState": (".simulation_manager", "SimulationState"),
|
||||
"SimulationStatus": (".simulation_manager", "SimulationStatus"),
|
||||
"SimulationConfigGenerator": (".simulation_config_generator", "SimulationConfigGenerator"),
|
||||
"SimulationParameters": (".simulation_config_generator", "SimulationParameters"),
|
||||
"AgentActivityConfig": (".simulation_config_generator", "AgentActivityConfig"),
|
||||
"TimeSimulationConfig": (".simulation_config_generator", "TimeSimulationConfig"),
|
||||
"EventConfig": (".simulation_config_generator", "EventConfig"),
|
||||
"PlatformConfig": (".simulation_config_generator", "PlatformConfig"),
|
||||
"SimulationRunner": (".simulation_runner", "SimulationRunner"),
|
||||
"SimulationRunState": (".simulation_runner", "SimulationRunState"),
|
||||
"RunnerStatus": (".simulation_runner", "RunnerStatus"),
|
||||
"AgentAction": (".simulation_runner", "AgentAction"),
|
||||
"RoundSummary": (".simulation_runner", "RoundSummary"),
|
||||
"ZepGraphMemoryUpdater": (".zep_graph_memory_updater", "ZepGraphMemoryUpdater"),
|
||||
"ZepGraphMemoryManager": (".zep_graph_memory_updater", "ZepGraphMemoryManager"),
|
||||
"AgentActivity": (".memory_activity", "AgentActivity"),
|
||||
"SimulationIPCClient": (".simulation_ipc", "SimulationIPCClient"),
|
||||
"SimulationIPCServer": (".simulation_ipc", "SimulationIPCServer"),
|
||||
"IPCCommand": (".simulation_ipc", "IPCCommand"),
|
||||
"IPCResponse": (".simulation_ipc", "IPCResponse"),
|
||||
"CommandType": (".simulation_ipc", "CommandType"),
|
||||
"CommandStatus": (".simulation_ipc", "CommandStatus"),
|
||||
}
|
||||
|
||||
|
||||
def __getattr__(name: str):
|
||||
try:
|
||||
module_name, attribute_name = _EXPORTS[name]
|
||||
except KeyError as exc:
|
||||
raise AttributeError(f"module {__name__!r} has no attribute {name!r}") from exc
|
||||
module = import_module(module_name, __name__)
|
||||
value = getattr(module, attribute_name)
|
||||
globals()[name] = value
|
||||
return value
|
||||
|
||||
|
||||
def __dir__() -> list[str]:
|
||||
return sorted(set(globals()) | set(_EXPORTS))
|
||||
|
||||
|
||||
__all__ = sorted(_EXPORTS)
|
||||
|
||||
|
||||
87
backend/app/services/artifact_store.py
Normal file
87
backend/app/services/artifact_store.py
Normal file
@@ -0,0 +1,87 @@
|
||||
"""Tenant-scoped artifact store abstraction.
|
||||
|
||||
Wraps filesystem artifact persistence behind a small interface so project files,
|
||||
simulation artifacts, and reports can later be moved to object storage without
|
||||
changing callers. All paths are resolved under a tenant directory and reject
|
||||
traversal/absolute components (fail-closed).
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
import re
|
||||
from typing import Optional
|
||||
|
||||
_SAFE_SEGMENT = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._-]{0,127}$")
|
||||
|
||||
|
||||
class ArtifactStore:
|
||||
"""Resolve tenant-scoped storage paths rooted under ``root``.
|
||||
|
||||
``path_for`` maps an owner/org and nested segments to an absolute path under
|
||||
``root/<organization>/<segments...>``. Path components are validated so a
|
||||
caller can never escape the configured root.
|
||||
"""
|
||||
|
||||
def __init__(self, root: str):
|
||||
if not isinstance(root, str) or not root.strip():
|
||||
raise ValueError("artifact_root_required")
|
||||
self.root = os.path.realpath(root)
|
||||
|
||||
def _validate_segments(self, segments):
|
||||
cleaned = []
|
||||
for segment in segments:
|
||||
if not isinstance(segment, str) or not _SAFE_SEGMENT.fullmatch(segment):
|
||||
raise ValueError("invalid_artifact_path")
|
||||
cleaned.append(segment)
|
||||
return cleaned
|
||||
|
||||
def path_for(self, *segments) -> str:
|
||||
"""Return a safe absolute path under the root for the given segments.
|
||||
|
||||
The first segment is treated as the tenant/owner scope; every segment
|
||||
must be a safe identifier (no separators, no dots-only, no traversal).
|
||||
"""
|
||||
if not segments:
|
||||
raise ValueError("artifact_path_required")
|
||||
cleaned = self._validate_segments(segments)
|
||||
candidate = os.path.realpath(os.path.join(self.root, *cleaned))
|
||||
if os.path.commonpath([self.root, candidate]) != self.root:
|
||||
raise ValueError("invalid_artifact_path")
|
||||
return candidate
|
||||
|
||||
def ensure_parent(self, path: str) -> None:
|
||||
parent = os.path.dirname(path)
|
||||
if parent:
|
||||
os.makedirs(parent, exist_ok=True)
|
||||
|
||||
def store_bytes(self, path: str, data: bytes) -> None:
|
||||
self.ensure_parent(path)
|
||||
with open(path, "wb") as handle:
|
||||
handle.write(data)
|
||||
|
||||
def read_bytes(self, path: str) -> Optional[bytes]:
|
||||
if not self.exists(path):
|
||||
return None
|
||||
with open(path, "rb") as handle:
|
||||
return handle.read()
|
||||
|
||||
def exists(self, path: str) -> bool:
|
||||
return os.path.isfile(path)
|
||||
|
||||
def delete(self, path: str) -> bool:
|
||||
if self.exists(path):
|
||||
os.remove(path)
|
||||
return True
|
||||
return False
|
||||
|
||||
|
||||
def default_artifact_store() -> ArtifactStore:
|
||||
"""Artifact store rooted at the configured upload root (filesystem backend).
|
||||
|
||||
Later this factory can return an object-storage backed store without
|
||||
changing callers.
|
||||
"""
|
||||
from ..config import Config
|
||||
|
||||
return ArtifactStore(Config.UPLOAD_FOLDER)
|
||||
63
backend/app/services/audit_service.py
Normal file
63
backend/app/services/audit_service.py
Normal file
@@ -0,0 +1,63 @@
|
||||
"""Durable, redacted audit event recording.
|
||||
|
||||
Audit details never store secrets, tokens, password hashes, API keys, or raw
|
||||
prompts — such keys are stripped before persisting. Events are tenant-scoped.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.operations import AuditLog
|
||||
|
||||
_SENSITIVE_KEYS = {"password", "token", "api_key", "secret", "authorization", "prompt"}
|
||||
|
||||
|
||||
class AuditService:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
@staticmethod
|
||||
def _redact(details) -> Optional[dict]:
|
||||
if not isinstance(details, dict):
|
||||
return None
|
||||
return {
|
||||
key: value
|
||||
for key, value in details.items()
|
||||
if key.lower() not in _SENSITIVE_KEYS
|
||||
}
|
||||
|
||||
def record(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
actor_user_id: Optional[str] = None,
|
||||
action: str,
|
||||
target_type: str,
|
||||
target_id: Optional[str] = None,
|
||||
details: Optional[dict] = None,
|
||||
) -> str:
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
raise ValueError("organization_id_required")
|
||||
entry = AuditLog(
|
||||
organization_id=organization_id,
|
||||
actor_user_id=actor_user_id,
|
||||
action=action,
|
||||
target_type=target_type,
|
||||
target_id=target_id,
|
||||
details=self._redact(details),
|
||||
)
|
||||
self.session.add(entry)
|
||||
self.session.flush()
|
||||
return entry.id
|
||||
|
||||
def list_for_organization(self, *, organization_id: str, limit: int = 100) -> list[AuditLog]:
|
||||
return (
|
||||
self.session.query(AuditLog)
|
||||
.filter(AuditLog.organization_id == organization_id)
|
||||
.order_by(AuditLog.created_at.desc())
|
||||
.limit(min(max(int(limit), 1), 1000))
|
||||
.all()
|
||||
)
|
||||
@@ -10,15 +10,15 @@ import threading
|
||||
from typing import Dict, Any, List, Optional, Callable
|
||||
from dataclasses import dataclass
|
||||
|
||||
from zep_cloud.client import Zep
|
||||
from zep_cloud import EpisodeData, EntityEdgeSourceTarget
|
||||
|
||||
from ..config import Config
|
||||
from ..models.task import TaskManager, TaskStatus
|
||||
from ..utils.zep_paging import fetch_all_nodes, fetch_all_edges
|
||||
from ..utils.logger import get_logger
|
||||
from .text_processor import TextProcessor
|
||||
from ..utils.locale import t, get_locale, set_locale
|
||||
|
||||
logger = get_logger('crowdsight.graph_builder')
|
||||
Zep = None
|
||||
|
||||
|
||||
@dataclass
|
||||
class GraphInfo:
|
||||
@@ -43,13 +43,28 @@ class GraphBuilderService:
|
||||
负责调用Zep API构建知识图谱
|
||||
"""
|
||||
|
||||
def __init__(self, api_key: Optional[str] = None):
|
||||
def __init__(
|
||||
self,
|
||||
api_key: Optional[str] = None,
|
||||
*,
|
||||
organization_id: Optional[str] = None,
|
||||
owner_user_id: Optional[str] = None,
|
||||
session_factory=None,
|
||||
):
|
||||
self.api_key = api_key or Config.ZEP_API_KEY
|
||||
if not self.api_key:
|
||||
raise ValueError("ZEP_API_KEY 未配置")
|
||||
|
||||
|
||||
global Zep
|
||||
if Zep is None:
|
||||
from zep_cloud.client import Zep as ZepClient
|
||||
|
||||
Zep = ZepClient
|
||||
|
||||
self.client = Zep(api_key=self.api_key)
|
||||
self.task_manager = TaskManager()
|
||||
self.organization_id = organization_id
|
||||
self.owner_user_id = owner_user_id
|
||||
self.task_manager = TaskManager(session_factory=session_factory)
|
||||
|
||||
def build_graph_async(
|
||||
self,
|
||||
@@ -78,6 +93,8 @@ class GraphBuilderService:
|
||||
task_id = self.task_manager.create_task(
|
||||
task_type="graph_build",
|
||||
metadata={
|
||||
"organization_id": self.organization_id,
|
||||
"owner_user_id": self.owner_user_id,
|
||||
"graph_name": graph_name,
|
||||
"chunk_size": chunk_size,
|
||||
"text_length": len(text),
|
||||
@@ -186,9 +203,12 @@ class GraphBuilderService:
|
||||
})
|
||||
|
||||
except Exception as e:
|
||||
import traceback
|
||||
error_msg = f"{str(e)}\n{traceback.format_exc()}"
|
||||
self.task_manager.fail_task(task_id, error_msg)
|
||||
logger.error(
|
||||
"Graph build failed: task_id=%s error_type=%s",
|
||||
task_id,
|
||||
type(e).__name__,
|
||||
)
|
||||
self.task_manager.fail_task(task_id, t('api.internalError'))
|
||||
|
||||
def create_graph(self, name: str) -> str:
|
||||
"""创建Zep图谱(公开方法)"""
|
||||
@@ -207,6 +227,7 @@ class GraphBuilderService:
|
||||
import warnings
|
||||
from typing import Optional
|
||||
from pydantic import Field
|
||||
from zep_cloud import EntityEdgeSourceTarget
|
||||
from zep_cloud.external_clients.ontology import EntityModel, EntityText, EdgeModel
|
||||
|
||||
# 抑制 Pydantic v2 关于 Field(default=None) 的警告
|
||||
@@ -315,6 +336,8 @@ class GraphBuilderService:
|
||||
)
|
||||
|
||||
# 构建episode数据
|
||||
from zep_cloud import EpisodeData
|
||||
|
||||
episodes = [
|
||||
EpisodeData(data=chunk, type="text")
|
||||
for chunk in batch_chunks
|
||||
@@ -402,6 +425,8 @@ class GraphBuilderService:
|
||||
|
||||
def _get_graph_info(self, graph_id: str) -> GraphInfo:
|
||||
"""获取图谱信息"""
|
||||
from ..utils.zep_paging import fetch_all_edges, fetch_all_nodes
|
||||
|
||||
# 获取节点(分页)
|
||||
nodes = fetch_all_nodes(self.client, graph_id)
|
||||
|
||||
@@ -433,6 +458,8 @@ class GraphBuilderService:
|
||||
Returns:
|
||||
包含nodes和edges的字典,包括时间信息、属性等详细数据
|
||||
"""
|
||||
from ..utils.zep_paging import fetch_all_edges, fetch_all_nodes
|
||||
|
||||
nodes = fetch_all_nodes(self.client, graph_id)
|
||||
edges = fetch_all_edges(self.client, graph_id)
|
||||
|
||||
|
||||
244
backend/app/services/idempotency.py
Normal file
244
backend/app/services/idempotency.py
Normal file
@@ -0,0 +1,244 @@
|
||||
"""Database-backed idempotency for cookie-authenticated mutations."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import json
|
||||
import re
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from functools import wraps
|
||||
from typing import Any
|
||||
|
||||
from flask import current_app, g, jsonify, make_response, request
|
||||
from sqlalchemy import delete, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.operations import IdempotencyRecord
|
||||
from ..utils.api_errors import ApiError
|
||||
from ..utils.locale import t
|
||||
|
||||
_KEY_RE = re.compile(r"^[A-Za-z0-9][A-Za-z0-9._:-]{0,127}$")
|
||||
|
||||
|
||||
class IdempotencyConflict(ValueError):
|
||||
"""The key was reused for a different request body."""
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class Reservation:
|
||||
record: IdempotencyRecord
|
||||
is_new: bool
|
||||
|
||||
|
||||
class IdempotencyService:
|
||||
DEFAULT_TTL_SECONDS = 24 * 60 * 60
|
||||
|
||||
def __init__(self, session: Session, *, ttl_seconds: int = DEFAULT_TTL_SECONDS):
|
||||
if ttl_seconds < 60 or ttl_seconds > 7 * 24 * 60 * 60:
|
||||
raise ValueError("invalid_idempotency_ttl")
|
||||
self.session = session
|
||||
self.ttl_seconds = ttl_seconds
|
||||
|
||||
@staticmethod
|
||||
def _request_hash(request_body: Any) -> str:
|
||||
try:
|
||||
encoded = json.dumps(
|
||||
request_body,
|
||||
ensure_ascii=False,
|
||||
sort_keys=True,
|
||||
separators=(",", ":"),
|
||||
).encode("utf-8")
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("invalid_idempotency_body") from exc
|
||||
return hashlib.sha256(encoded).hexdigest()
|
||||
|
||||
@staticmethod
|
||||
def validate_key(key: str) -> str:
|
||||
if not isinstance(key, str) or not _KEY_RE.fullmatch(key):
|
||||
raise ValueError("invalid_idempotency_key")
|
||||
return key
|
||||
|
||||
def reserve(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
user_id: str,
|
||||
key: str,
|
||||
request_body: Any,
|
||||
) -> Reservation:
|
||||
if not organization_id or not user_id:
|
||||
raise ValueError("invalid_idempotency_scope")
|
||||
normalized_key = self.validate_key(key)
|
||||
request_hash = self._request_hash(request_body)
|
||||
now = datetime.now(timezone.utc)
|
||||
existing = self.session.scalar(
|
||||
select(IdempotencyRecord)
|
||||
.where(
|
||||
IdempotencyRecord.organization_id == organization_id,
|
||||
IdempotencyRecord.user_id == user_id,
|
||||
IdempotencyRecord.key == normalized_key,
|
||||
)
|
||||
.with_for_update()
|
||||
)
|
||||
if existing is not None:
|
||||
expires_at = existing.expires_at
|
||||
if expires_at.tzinfo is None:
|
||||
expires_at = expires_at.replace(tzinfo=timezone.utc)
|
||||
else:
|
||||
expires_at = None
|
||||
if existing is not None and expires_at is not None and expires_at <= now:
|
||||
self.session.execute(delete(IdempotencyRecord).where(IdempotencyRecord.id == existing.id))
|
||||
self.session.flush()
|
||||
existing = None
|
||||
if existing is not None:
|
||||
if existing.request_hash != request_hash:
|
||||
raise IdempotencyConflict("idempotency_key_reused")
|
||||
return Reservation(record=existing, is_new=False)
|
||||
|
||||
record = IdempotencyRecord(
|
||||
organization_id=organization_id,
|
||||
user_id=user_id,
|
||||
key=normalized_key,
|
||||
request_hash=request_hash,
|
||||
status="reserved",
|
||||
expires_at=now + timedelta(seconds=self.ttl_seconds),
|
||||
)
|
||||
self.session.add(record)
|
||||
self.session.flush()
|
||||
return Reservation(record=record, is_new=True)
|
||||
|
||||
def complete(self, record: IdempotencyRecord, *, status_code: int, body: dict | list) -> None:
|
||||
if status_code < 100 or status_code > 599:
|
||||
raise ValueError("invalid_response_status")
|
||||
record.status = "completed"
|
||||
record.response_status = status_code
|
||||
record.response_body = body
|
||||
record.completed_at = datetime.now(timezone.utc)
|
||||
self.session.flush()
|
||||
|
||||
def fail(self, record: IdempotencyRecord, *, status_code: int, body: dict | list) -> None:
|
||||
if status_code < 400 or status_code > 599:
|
||||
raise ValueError("invalid_failure_status")
|
||||
record.status = "failed"
|
||||
record.response_status = status_code
|
||||
record.response_body = body
|
||||
record.completed_at = datetime.now(timezone.utc)
|
||||
self.session.flush()
|
||||
|
||||
|
||||
def _stream_sha256(stream) -> str:
|
||||
try:
|
||||
start_position = stream.tell()
|
||||
except (AttributeError, OSError, ValueError) as exc:
|
||||
raise ValueError("invalid_idempotency_body") from exc
|
||||
|
||||
digest = hashlib.sha256()
|
||||
try:
|
||||
while True:
|
||||
chunk = stream.read(1024 * 1024)
|
||||
if not chunk:
|
||||
break
|
||||
digest.update(chunk)
|
||||
except (AttributeError, OSError, TypeError, ValueError) as exc:
|
||||
raise ValueError("invalid_idempotency_body") from exc
|
||||
finally:
|
||||
try:
|
||||
stream.seek(start_position)
|
||||
except (AttributeError, OSError, ValueError) as exc:
|
||||
raise ValueError("invalid_idempotency_body") from exc
|
||||
return digest.hexdigest()
|
||||
|
||||
|
||||
def _request_fingerprint_payload() -> Any:
|
||||
if request.is_json:
|
||||
body = request.get_json(silent=True)
|
||||
elif request.form or request.files:
|
||||
files = [
|
||||
{
|
||||
"field": field,
|
||||
"filename": file.filename or "",
|
||||
"content_type": file.content_type or "",
|
||||
"size": request.content_length or 0,
|
||||
"content_sha256": _stream_sha256(file.stream),
|
||||
}
|
||||
for field, files in request.files.lists()
|
||||
for file in files
|
||||
]
|
||||
body = {"form": request.form.to_dict(flat=False), "files": files}
|
||||
else:
|
||||
body = {}
|
||||
return {
|
||||
"method": request.method,
|
||||
"path": request.path,
|
||||
"query": request.args.to_dict(flat=False),
|
||||
"body": body,
|
||||
}
|
||||
|
||||
|
||||
def idempotent(view):
|
||||
"""Require and persist an ``Idempotency-Key`` for a JSON mutation."""
|
||||
|
||||
@wraps(view)
|
||||
def wrapped(*args, **kwargs):
|
||||
key = request.headers.get("Idempotency-Key", "")
|
||||
if not key:
|
||||
raise ApiError("idempotency_required", 400, "api.idempotencyRequired")
|
||||
session = getattr(g, "db_session", None)
|
||||
managed_session = False
|
||||
if session is None:
|
||||
factory = current_app.extensions.get("crowdsight_session_factory")
|
||||
if not callable(factory):
|
||||
raise ApiError("auth_unavailable", 503, "api.internalError")
|
||||
session = factory()
|
||||
managed_session = True
|
||||
context = getattr(g, "auth_context", None)
|
||||
if context is None:
|
||||
if managed_session:
|
||||
session.close()
|
||||
raise ApiError("auth_unavailable", 503, "api.internalError")
|
||||
try:
|
||||
try:
|
||||
reservation = IdempotencyService(session).reserve(
|
||||
organization_id=context.organization.id,
|
||||
user_id=context.user.id,
|
||||
key=key,
|
||||
request_body=_request_fingerprint_payload(),
|
||||
)
|
||||
except IdempotencyConflict as exc:
|
||||
raise ApiError("idempotency_key_reused", 409, "api.idempotencyConflict") from exc
|
||||
except ValueError as exc:
|
||||
code = str(exc)
|
||||
if code == "invalid_idempotency_key":
|
||||
raise ApiError(code, 400, "api.idempotencyInvalid") from exc
|
||||
raise ApiError("invalid_request", 400, "api.requestError") from exc
|
||||
|
||||
if not reservation.is_new:
|
||||
record = reservation.record
|
||||
if record.status == "completed" and record.response_body is not None:
|
||||
return jsonify(record.response_body), record.response_status or 200
|
||||
if record.status == "failed" and record.response_body is not None:
|
||||
return jsonify(record.response_body), record.response_status or 500
|
||||
raise ApiError("idempotency_in_progress", 409, "api.idempotencyInProgress")
|
||||
|
||||
response = make_response(view(*args, **kwargs))
|
||||
body = response.get_json(silent=True)
|
||||
if not isinstance(body, (dict, list)):
|
||||
raise ApiError("idempotency_response_invalid", 500, "api.internalError")
|
||||
IdempotencyService(session).complete(
|
||||
reservation.record,
|
||||
status_code=response.status_code,
|
||||
body=body,
|
||||
)
|
||||
if managed_session:
|
||||
session.commit()
|
||||
return response
|
||||
except Exception:
|
||||
if managed_session:
|
||||
session.rollback()
|
||||
raise
|
||||
finally:
|
||||
if managed_session:
|
||||
session.close()
|
||||
|
||||
return wrapped
|
||||
260
backend/app/services/identity.py
Normal file
260
backend/app/services/identity.py
Normal file
@@ -0,0 +1,260 @@
|
||||
"""Tenant-scoped identity repository and password service."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from dataclasses import dataclass
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from argon2 import PasswordHasher
|
||||
from argon2.exceptions import InvalidHashError, VerificationError, VerifyMismatchError
|
||||
from sqlalchemy import select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.saas import AuthSession, Membership, Organization, User
|
||||
from ..security.policy import Role
|
||||
|
||||
|
||||
class PasswordService:
|
||||
"""Argon2id password hashing with no plaintext fallback."""
|
||||
|
||||
_hasher = PasswordHasher()
|
||||
|
||||
@classmethod
|
||||
def hash_password(cls, password: str) -> str:
|
||||
if not isinstance(password, str) or len(password) < 12:
|
||||
raise ValueError("password_too_short")
|
||||
return cls._hasher.hash(password)
|
||||
|
||||
@classmethod
|
||||
def verify_password(cls, password_hash: str, password: str) -> bool:
|
||||
if not isinstance(password_hash, str) or not password_hash:
|
||||
return False
|
||||
try:
|
||||
return cls._hasher.verify(password_hash, password)
|
||||
except (VerifyMismatchError, VerificationError, InvalidHashError):
|
||||
return False
|
||||
|
||||
|
||||
class IdentityRepository:
|
||||
"""Flush-only repository; the caller owns transaction boundaries."""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
@staticmethod
|
||||
def normalize_email(email: str) -> str:
|
||||
if not isinstance(email, str):
|
||||
raise ValueError("invalid_email")
|
||||
normalized = email.strip().casefold()
|
||||
if "@" not in normalized or normalized.startswith("@") or normalized.endswith("@"):
|
||||
raise ValueError("invalid_email")
|
||||
return normalized
|
||||
|
||||
@staticmethod
|
||||
def normalize_slug(slug: str) -> str:
|
||||
if not isinstance(slug, str):
|
||||
raise ValueError("invalid_slug")
|
||||
normalized = slug.strip().casefold()
|
||||
if not normalized or any(char not in "abcdefghijklmnopqrstuvwxyz0123456789-" for char in normalized):
|
||||
raise ValueError("invalid_slug")
|
||||
return normalized
|
||||
|
||||
def create_organization(self, *, name: str, slug: str) -> Organization:
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("invalid_organization_name")
|
||||
organization = Organization(name=name.strip(), slug=self.normalize_slug(slug))
|
||||
self.session.add(organization)
|
||||
self.session.flush()
|
||||
return organization
|
||||
|
||||
def create_user(self, *, email: str, password_hash: str | None = None) -> User:
|
||||
user = User(
|
||||
email_normalized=self.normalize_email(email),
|
||||
password_hash=password_hash or "!invite_pending",
|
||||
)
|
||||
self.session.add(user)
|
||||
self.session.flush()
|
||||
return user
|
||||
|
||||
def create_membership(self, user_id: str, organization_id: str, role: Role | str) -> Membership:
|
||||
try:
|
||||
normalized_role = role if isinstance(role, Role) else Role(role)
|
||||
except (TypeError, ValueError) as exc:
|
||||
raise ValueError("invalid_role") from exc
|
||||
|
||||
membership = Membership(
|
||||
user_id=user_id,
|
||||
organization_id=organization_id,
|
||||
role=normalized_role,
|
||||
)
|
||||
self.session.add(membership)
|
||||
self.session.flush()
|
||||
return membership
|
||||
|
||||
def get_organization(self, organization_id: str) -> Organization | None:
|
||||
return self.session.scalar(
|
||||
select(Organization).where(Organization.id == organization_id)
|
||||
)
|
||||
|
||||
def get_organization_by_slug(self, slug: str) -> Organization | None:
|
||||
normalized = self.normalize_slug(slug)
|
||||
return self.session.scalar(
|
||||
select(Organization).where(Organization.slug == normalized)
|
||||
)
|
||||
|
||||
def list_active_memberships(self, user_id: str) -> list[tuple[Membership, Organization]]:
|
||||
return list(
|
||||
self.session.execute(
|
||||
select(Membership, Organization)
|
||||
.join(Organization, Organization.id == Membership.organization_id)
|
||||
.where(
|
||||
Membership.user_id == user_id,
|
||||
Membership.status == "active",
|
||||
Organization.status == "active",
|
||||
)
|
||||
.order_by(Organization.slug)
|
||||
).all()
|
||||
)
|
||||
|
||||
def get_user_by_email(self, email: str) -> User | None:
|
||||
normalized = self.normalize_email(email)
|
||||
return self.session.scalar(
|
||||
select(User).where(User.email_normalized == normalized)
|
||||
)
|
||||
|
||||
def get_user_for_org(self, user_id: str, organization_id: str) -> User | None:
|
||||
"""Tenant-scoped read; wrong-org IDs return None without disclosure."""
|
||||
return self.session.scalar(
|
||||
select(User)
|
||||
.join(Membership, Membership.user_id == User.id)
|
||||
.where(
|
||||
User.id == user_id,
|
||||
Membership.organization_id == organization_id,
|
||||
Membership.status == "active",
|
||||
User.status == "active",
|
||||
)
|
||||
)
|
||||
|
||||
def list_users_with_memberships(self, organization_id: str) -> list[tuple[User, Membership]]:
|
||||
return list(
|
||||
self.session.execute(
|
||||
select(User, Membership)
|
||||
.join(Membership, Membership.user_id == User.id)
|
||||
.where(
|
||||
Membership.organization_id == organization_id,
|
||||
Membership.status == "active",
|
||||
User.status == "active",
|
||||
)
|
||||
.order_by(User.email_normalized)
|
||||
).all()
|
||||
)
|
||||
|
||||
def list_users(self, organization_id: str) -> list[User]:
|
||||
"""Return only active users with active membership in the tenant."""
|
||||
return list(
|
||||
self.session.scalars(
|
||||
select(User)
|
||||
.join(Membership, Membership.user_id == User.id)
|
||||
.where(
|
||||
Membership.organization_id == organization_id,
|
||||
Membership.status == "active",
|
||||
User.status == "active",
|
||||
)
|
||||
.order_by(User.email_normalized)
|
||||
)
|
||||
)
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class SessionContext:
|
||||
user: User
|
||||
membership: Membership
|
||||
organization: Organization
|
||||
|
||||
|
||||
class SessionService:
|
||||
"""Opaque, revocable session token operations; caller owns commits."""
|
||||
|
||||
DEFAULT_TTL_SECONDS = 60 * 60 * 12
|
||||
|
||||
@staticmethod
|
||||
def _hash_token(raw_token: str) -> str:
|
||||
return hashlib.sha256(raw_token.encode("utf-8")).hexdigest()
|
||||
|
||||
@classmethod
|
||||
def create(
|
||||
cls,
|
||||
session: Session,
|
||||
user: User,
|
||||
membership_id: str,
|
||||
ttl_seconds: int | None = None,
|
||||
) -> tuple[str, AuthSession]:
|
||||
if not isinstance(membership_id, str) or not membership_id:
|
||||
raise ValueError("invalid_membership")
|
||||
ttl = ttl_seconds or cls.DEFAULT_TTL_SECONDS
|
||||
if ttl < 60:
|
||||
raise ValueError("invalid_session_ttl")
|
||||
|
||||
membership = session.scalar(
|
||||
select(Membership).where(
|
||||
Membership.id == membership_id,
|
||||
Membership.user_id == user.id,
|
||||
Membership.status == "active",
|
||||
)
|
||||
)
|
||||
if membership is None:
|
||||
raise ValueError("invalid_membership")
|
||||
|
||||
raw_token = secrets.token_urlsafe(32)
|
||||
stored = AuthSession(
|
||||
user_id=user.id,
|
||||
membership_id=membership_id,
|
||||
token_hash=cls._hash_token(raw_token),
|
||||
auth_version=user.auth_version,
|
||||
expires_at=datetime.now(timezone.utc) + timedelta(seconds=ttl),
|
||||
)
|
||||
session.add(stored)
|
||||
session.flush()
|
||||
return raw_token, stored
|
||||
|
||||
@classmethod
|
||||
def resolve(cls, session: Session, raw_token: str) -> SessionContext | None:
|
||||
if not isinstance(raw_token, str) or not raw_token:
|
||||
return None
|
||||
|
||||
now = datetime.now(timezone.utc)
|
||||
row = session.execute(
|
||||
select(AuthSession, User, Membership, Organization)
|
||||
.join(User, User.id == AuthSession.user_id)
|
||||
.join(Membership, Membership.id == AuthSession.membership_id)
|
||||
.join(Organization, Organization.id == Membership.organization_id)
|
||||
.where(
|
||||
AuthSession.token_hash == cls._hash_token(raw_token),
|
||||
AuthSession.revoked_at.is_(None),
|
||||
AuthSession.expires_at > now,
|
||||
AuthSession.auth_version == User.auth_version,
|
||||
User.status == "active",
|
||||
Membership.status == "active",
|
||||
Organization.status == "active",
|
||||
)
|
||||
).first()
|
||||
if row is None:
|
||||
return None
|
||||
|
||||
stored, user, membership, organization = row
|
||||
stored.last_seen_at = now
|
||||
return SessionContext(user=user, membership=membership, organization=organization)
|
||||
|
||||
@classmethod
|
||||
def revoke(cls, session: Session, raw_token: str) -> bool:
|
||||
if not isinstance(raw_token, str) or not raw_token:
|
||||
return False
|
||||
stored = session.scalar(
|
||||
select(AuthSession).where(AuthSession.token_hash == cls._hash_token(raw_token))
|
||||
)
|
||||
if stored is None or stored.revoked_at is not None:
|
||||
return False
|
||||
stored.revoked_at = datetime.now(timezone.utc)
|
||||
session.flush()
|
||||
return True
|
||||
82
backend/app/services/job_queue.py
Normal file
82
backend/app/services/job_queue.py
Normal file
@@ -0,0 +1,82 @@
|
||||
"""Durable job queue: claim/complete/fail lifecycle over the ``jobs`` table.
|
||||
|
||||
This is the portable core a worker topology builds on; it does not require a
|
||||
broker. For PostgreSQL production this should issue ``SELECT ... FOR UPDATE``
|
||||
plus ``UPDATE ... WHERE status='queued'`` to claim atomically; the SQLite
|
||||
fallback below uses a synchronized in-process write for local/tests.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timezone
|
||||
from typing import Callable, Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.operations import Job, JobStatus
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class JobQueue:
|
||||
"""Flush-only queue: caller owns the transaction boundary."""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
self._handlers: dict[str, Callable] = {}
|
||||
|
||||
def register_handler(self, operation: str, handler) -> None:
|
||||
"""Register a callable handler for an operation (worker plugin point)."""
|
||||
self._handlers[operation] = handler
|
||||
|
||||
def dispatch(self, job: Job, *, payload=None):
|
||||
"""Invoke the registered handler for ``job.operation``.
|
||||
|
||||
Returns the handler result. Raises ``ValueError`` when no handler is
|
||||
registered so the worker can fail the job.
|
||||
"""
|
||||
handler = self._handlers.get(job.operation)
|
||||
if handler is None:
|
||||
raise ValueError(f"no_handler_for_operation: {job.operation}")
|
||||
return handler(payload, job)
|
||||
|
||||
def claim_next_job(
|
||||
self, *, worker_id: str = "worker", organization_id: Optional[str] = None
|
||||
) -> Optional[Job]:
|
||||
"""Claim the next queued job for a worker.
|
||||
|
||||
For the given optional organization scope, atomically flip one queued
|
||||
job to ``running`` and return it; returns None when nothing is claimable.
|
||||
"""
|
||||
query = self.session.query(Job).filter(Job.status == JobStatus.QUEUED.value)
|
||||
if organization_id is not None:
|
||||
query = query.filter(Job.organization_id == organization_id)
|
||||
job = query.order_by(Job.created_at.asc()).first()
|
||||
if job is None:
|
||||
return None
|
||||
job.status = JobStatus.RUNNING.value
|
||||
job.message = f"claimed by {worker_id}"
|
||||
self.session.flush()
|
||||
return job
|
||||
|
||||
def complete_job(self, job_id: str, result=None, message: str = "") -> None:
|
||||
job = self.session.get(Job, job_id)
|
||||
if job is None:
|
||||
raise ValueError("job_not_found")
|
||||
job.status = JobStatus.SUCCEEDED.value
|
||||
job.result = result
|
||||
job.message = message or job.message
|
||||
job.finished_at = _utc_now()
|
||||
self.session.flush()
|
||||
|
||||
def fail_job(self, job_id: str, error_code: str, message: str = "") -> None:
|
||||
job = self.session.get(Job, job_id)
|
||||
if job is None:
|
||||
raise ValueError("job_not_found")
|
||||
job.status = JobStatus.FAILED.value
|
||||
job.error_code = error_code
|
||||
job.message = message or job.message
|
||||
job.finished_at = _utc_now()
|
||||
self.session.flush()
|
||||
197
backend/app/services/local_graph_builder.py
Normal file
197
backend/app/services/local_graph_builder.py
Normal file
@@ -0,0 +1,197 @@
|
||||
"""Tenant-scoped graph builder backed by the local memory repository."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Any, Callable, Iterable
|
||||
from uuid import uuid4
|
||||
|
||||
from .memory_repository import SqlAlchemyMemoryRepository
|
||||
from .memory_service import MemoryExtractionService
|
||||
from ..utils.locale import get_locale
|
||||
from ..utils.language_policy import normalize_locale
|
||||
from ..utils.logger import get_logger
|
||||
|
||||
logger = get_logger("crowdsight.local_graph_builder")
|
||||
|
||||
|
||||
ProgressCallback = Callable[[str, float], None]
|
||||
|
||||
|
||||
class LocalGraphBuilderService:
|
||||
"""Build a graph synchronously into local durable memory.
|
||||
|
||||
The public methods intentionally mirror ``GraphBuilderService`` so the API
|
||||
can switch storage backends without giving either backend authorization
|
||||
authority. Every repository instance is created with the request actor's
|
||||
organization and the graph being processed.
|
||||
"""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
session_factory,
|
||||
*,
|
||||
organization_id: str,
|
||||
project_id: str,
|
||||
extraction_client=None,
|
||||
extraction_service: MemoryExtractionService | None = None,
|
||||
language: str | None = None,
|
||||
):
|
||||
if not callable(session_factory):
|
||||
raise ValueError("memory_session_factory_required")
|
||||
if not isinstance(organization_id, str) or not organization_id.strip():
|
||||
raise ValueError("memory_builder_organization_required")
|
||||
if not isinstance(project_id, str) or not project_id.strip():
|
||||
raise ValueError("memory_builder_project_required")
|
||||
self.session_factory = session_factory
|
||||
self.organization_id = organization_id
|
||||
self.project_id = project_id
|
||||
self.language = normalize_locale(language or get_locale())
|
||||
self.extraction_service = (
|
||||
extraction_service
|
||||
if extraction_service is not None
|
||||
else MemoryExtractionService(extraction_client) if extraction_client is not None else None
|
||||
)
|
||||
self._ontology_by_graph: dict[str, dict[str, Any]] = {}
|
||||
|
||||
def _repository(self, session, graph_id: str) -> SqlAlchemyMemoryRepository:
|
||||
return SqlAlchemyMemoryRepository(
|
||||
session,
|
||||
organization_id=self.organization_id,
|
||||
graph_id=graph_id,
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _validate_graph_id(graph_id: str) -> str:
|
||||
if not isinstance(graph_id, str) or not graph_id.strip():
|
||||
raise ValueError("invalid_graph_id")
|
||||
return graph_id.strip()
|
||||
|
||||
def create_graph(self, name: str = "") -> str:
|
||||
"""Create a new graph and return its durable ID."""
|
||||
graph_id = f"crowdsight_{uuid4().hex[:24]}"
|
||||
with self.session_factory() as session:
|
||||
repository = self._repository(session, graph_id)
|
||||
repository.create_graph(project_id=self.project_id)
|
||||
session.commit()
|
||||
logger.info("Created local memory graph %s", graph_id)
|
||||
return graph_id
|
||||
|
||||
def set_ontology(self, graph_id: str, ontology: dict[str, Any]) -> None:
|
||||
graph_id = self._validate_graph_id(graph_id)
|
||||
if not isinstance(ontology, dict):
|
||||
raise ValueError("invalid_ontology")
|
||||
with self.session_factory() as session:
|
||||
repository = self._repository(session, graph_id)
|
||||
repository.update_graph_ontology(ontology)
|
||||
session.commit()
|
||||
self._ontology_by_graph[graph_id] = dict(ontology)
|
||||
|
||||
def add_text_batches(
|
||||
self,
|
||||
graph_id: str,
|
||||
text_batches: Iterable[str],
|
||||
batch_size: int = 3,
|
||||
progress_callback: ProgressCallback | None = None,
|
||||
) -> list[str]:
|
||||
graph_id = self._validate_graph_id(graph_id)
|
||||
extraction_service = self.extraction_service
|
||||
if extraction_service is None:
|
||||
raise ValueError("memory_extraction_client_required")
|
||||
if isinstance(batch_size, bool) or not isinstance(batch_size, int) or batch_size < 1:
|
||||
raise ValueError("invalid_batch_size")
|
||||
chunks = list(text_batches)
|
||||
if any(not isinstance(chunk, str) or not chunk.strip() for chunk in chunks):
|
||||
raise ValueError("invalid_text_batch")
|
||||
|
||||
episode_ids: list[str] = []
|
||||
total = len(chunks)
|
||||
for index, chunk in enumerate(chunks):
|
||||
with self.session_factory() as session:
|
||||
repository = self._repository(session, graph_id)
|
||||
graph = repository.get_graph()
|
||||
ontology = dict(self._ontology_by_graph.get(graph_id) or graph.ontology)
|
||||
result = extraction_service.extract(
|
||||
language=self.language,
|
||||
ontology=ontology,
|
||||
episode_text=chunk,
|
||||
)
|
||||
ingest_result = extraction_service.persist(
|
||||
repository,
|
||||
result,
|
||||
source_type="text",
|
||||
source_ref=f"episode_{index}",
|
||||
episode_text=chunk,
|
||||
)
|
||||
episode = repository.get_episode(source_type="text", source_ref=f"episode_{index}")
|
||||
session.commit()
|
||||
if episode is None:
|
||||
raise RuntimeError("local_episode_persist_failed")
|
||||
episode_ids.append(episode.id)
|
||||
if progress_callback:
|
||||
progress_callback(
|
||||
f"Processed local memory chunk {index + 1}/{total}",
|
||||
(index + 1) / total if total else 1.0,
|
||||
)
|
||||
logger.debug(
|
||||
"Processed local graph chunk %s: entities=%s edges=%s",
|
||||
index,
|
||||
ingest_result.entity_count,
|
||||
ingest_result.edge_count,
|
||||
)
|
||||
return episode_ids
|
||||
|
||||
def _wait_for_episodes(
|
||||
self,
|
||||
episode_ids: list[str],
|
||||
progress_callback: ProgressCallback | None = None,
|
||||
) -> None:
|
||||
"""Local extraction is committed synchronously; verify IDs instead of polling Zep."""
|
||||
if not isinstance(episode_ids, list):
|
||||
raise ValueError("invalid_episode_ids")
|
||||
if progress_callback:
|
||||
progress_callback("Local memory processing complete", 1.0)
|
||||
|
||||
def get_graph_data(self, graph_id: str) -> dict[str, Any]:
|
||||
graph_id = self._validate_graph_id(graph_id)
|
||||
with self.session_factory() as session:
|
||||
repository = self._repository(session, graph_id)
|
||||
graph = repository.get_graph()
|
||||
nodes = repository.list_nodes(limit=10_000)
|
||||
edges = repository.list_edges(limit=20_000)
|
||||
return {
|
||||
"graph_id": graph.id,
|
||||
"node_count": len(nodes),
|
||||
"edge_count": len(edges),
|
||||
"nodes": [
|
||||
{
|
||||
"uuid": node.id,
|
||||
"name": node.canonical_name,
|
||||
"labels": list(node.labels or []),
|
||||
"summary": node.summary or "",
|
||||
"attributes": dict(node.attributes or {}),
|
||||
}
|
||||
for node in nodes
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"uuid": edge.id,
|
||||
"name": edge.relation,
|
||||
"fact": edge.fact,
|
||||
"source_node_uuid": edge.source_node_id,
|
||||
"target_node_uuid": edge.target_node_id,
|
||||
"attributes": dict(edge.attributes or {}),
|
||||
}
|
||||
for edge in edges
|
||||
],
|
||||
}
|
||||
|
||||
def delete_graph(self, graph_id: str) -> None:
|
||||
"""Delete a graph only when it belongs to this organization scope."""
|
||||
graph_id = self._validate_graph_id(graph_id)
|
||||
with self.session_factory() as session:
|
||||
repository = self._repository(session, graph_id)
|
||||
graph = repository.get_graph()
|
||||
session.delete(graph)
|
||||
session.commit()
|
||||
self._ontology_by_graph.pop(graph_id, None)
|
||||
logger.info("Deleted local memory graph %s", graph_id)
|
||||
177
backend/app/services/local_graph_memory_updater.py
Normal file
177
backend/app/services/local_graph_memory_updater.py
Normal file
@@ -0,0 +1,177 @@
|
||||
"""Tenant-scoped local persistence for simulation activities.
|
||||
|
||||
This module intentionally does not import or construct a Zep client. It keeps
|
||||
runtime activity updates behind the same small manager interface used by the
|
||||
legacy runner, while storing each activity as a durable local memory episode.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import threading
|
||||
from datetime import datetime
|
||||
from typing import Any, Dict, Optional, cast
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from .memory_repository import SqlAlchemyMemoryRepository
|
||||
from ..utils.logger import get_logger
|
||||
|
||||
logger = get_logger("crowdsight.local_graph_memory_updater")
|
||||
|
||||
|
||||
class LocalGraphMemoryUpdater:
|
||||
"""Persist simulation activities in a tenant-scoped local graph."""
|
||||
|
||||
def __init__(
|
||||
self,
|
||||
simulation_id: str,
|
||||
graph_id: str,
|
||||
organization_id: str,
|
||||
session_factory,
|
||||
):
|
||||
for name, value in (
|
||||
("simulation_id", simulation_id),
|
||||
("graph_id", graph_id),
|
||||
("organization_id", organization_id),
|
||||
):
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError(f"local_memory_{name}_required")
|
||||
if not callable(session_factory):
|
||||
raise ValueError("memory_session_factory_required")
|
||||
|
||||
self.simulation_id = simulation_id
|
||||
self.graph_id = graph_id
|
||||
self.organization_id = organization_id
|
||||
self.session_factory = session_factory
|
||||
self._running = False
|
||||
self._sequence = 0
|
||||
self._lock = threading.Lock()
|
||||
self._total_activities = 0
|
||||
self._total_sent = 0
|
||||
self._total_items_sent = 0
|
||||
self._failed_count = 0
|
||||
self._skipped_count = 0
|
||||
|
||||
def start(self) -> None:
|
||||
self._running = True
|
||||
|
||||
def stop(self) -> None:
|
||||
self._running = False
|
||||
|
||||
def add_activity(self, activity: Any) -> None:
|
||||
if not callable(getattr(activity, "to_episode_text", None)):
|
||||
raise ValueError("invalid_agent_activity")
|
||||
if activity.action_type == "DO_NOTHING":
|
||||
self._skipped_count += 1
|
||||
return
|
||||
|
||||
with self._lock:
|
||||
self._sequence += 1
|
||||
sequence = self._sequence
|
||||
self._total_activities += 1
|
||||
|
||||
try:
|
||||
episode_text = activity.to_episode_text()
|
||||
source_ref = (
|
||||
f"{self.simulation_id}:{activity.platform}:{activity.round_num}:"
|
||||
f"{activity.agent_id}:{sequence}"
|
||||
)
|
||||
session: Session = cast(Session, self.session_factory())
|
||||
try:
|
||||
repository = SqlAlchemyMemoryRepository(
|
||||
session,
|
||||
organization_id=self.organization_id,
|
||||
graph_id=self.graph_id,
|
||||
)
|
||||
repository.add_episode(
|
||||
source_type="simulation_activity",
|
||||
source_ref=source_ref,
|
||||
normalized_text=episode_text,
|
||||
summary=episode_text,
|
||||
extractor_version="runtime-v1",
|
||||
)
|
||||
session.commit()
|
||||
finally:
|
||||
session.close()
|
||||
self._total_sent += 1
|
||||
self._total_items_sent += 1
|
||||
except Exception:
|
||||
self._failed_count += 1
|
||||
raise
|
||||
|
||||
def add_activity_from_dict(self, data: Dict[str, Any], platform: str) -> None:
|
||||
if "event_type" in data:
|
||||
return
|
||||
from .memory_activity import AgentActivity
|
||||
|
||||
self.add_activity(
|
||||
AgentActivity(
|
||||
platform=platform,
|
||||
agent_id=data.get("agent_id", 0),
|
||||
agent_name=data.get("agent_name", ""),
|
||||
action_type=data.get("action_type", ""),
|
||||
action_args=data.get("action_args", {}),
|
||||
round_num=data.get("round", 0),
|
||||
timestamp=data.get("timestamp", datetime.now().isoformat()),
|
||||
)
|
||||
)
|
||||
|
||||
def get_stats(self) -> Dict[str, Any]:
|
||||
return {
|
||||
"total_activities": self._total_activities,
|
||||
"batches_sent": self._total_sent,
|
||||
"items_sent": self._total_items_sent,
|
||||
"failed_count": self._failed_count,
|
||||
"skipped_count": self._skipped_count,
|
||||
"queue_size": 0,
|
||||
"buffer_sizes": {},
|
||||
"running": self._running,
|
||||
}
|
||||
|
||||
|
||||
class LocalGraphMemoryManager:
|
||||
"""Manage one local updater per simulation."""
|
||||
|
||||
_updaters: Dict[str, LocalGraphMemoryUpdater] = {}
|
||||
_lock = threading.Lock()
|
||||
|
||||
@classmethod
|
||||
def create_updater(
|
||||
cls,
|
||||
simulation_id: str,
|
||||
graph_id: str,
|
||||
*,
|
||||
organization_id: str,
|
||||
session_factory,
|
||||
) -> LocalGraphMemoryUpdater:
|
||||
with cls._lock:
|
||||
existing = cls._updaters.get(simulation_id)
|
||||
if existing is not None:
|
||||
existing.stop()
|
||||
updater = LocalGraphMemoryUpdater(
|
||||
simulation_id=simulation_id,
|
||||
graph_id=graph_id,
|
||||
organization_id=organization_id,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
updater.start()
|
||||
cls._updaters[simulation_id] = updater
|
||||
return updater
|
||||
|
||||
@classmethod
|
||||
def get_updater(cls, simulation_id: str) -> Optional[LocalGraphMemoryUpdater]:
|
||||
return cls._updaters.get(simulation_id)
|
||||
|
||||
@classmethod
|
||||
def stop_updater(cls, simulation_id: str) -> None:
|
||||
with cls._lock:
|
||||
updater = cls._updaters.pop(simulation_id, None)
|
||||
if updater is not None:
|
||||
updater.stop()
|
||||
|
||||
@classmethod
|
||||
def stop_all(cls) -> None:
|
||||
with cls._lock:
|
||||
simulation_ids = list(cls._updaters)
|
||||
for simulation_id in simulation_ids:
|
||||
cls.stop_updater(simulation_id)
|
||||
184
backend/app/services/memory_activity.py
Normal file
184
backend/app/services/memory_activity.py
Normal file
@@ -0,0 +1,184 @@
|
||||
"""Shared simulation activity contract for local and Zep memory backends."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass
|
||||
from typing import Any, Dict
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentActivity:
|
||||
"""Agent活动记录"""
|
||||
platform: str # twitter / reddit
|
||||
agent_id: int
|
||||
agent_name: str
|
||||
action_type: str # CREATE_POST, LIKE_POST, etc.
|
||||
action_args: Dict[str, Any]
|
||||
round_num: int
|
||||
timestamp: str
|
||||
|
||||
def to_episode_text(self) -> str:
|
||||
"""
|
||||
将活动转换为可以发送给Zep的文本描述
|
||||
|
||||
采用自然语言描述格式,让Zep能够从中提取实体和关系
|
||||
不添加模拟相关的前缀,避免误导图谱更新
|
||||
"""
|
||||
# 根据不同的动作类型生成不同的描述
|
||||
action_descriptions = {
|
||||
"CREATE_POST": self._describe_create_post,
|
||||
"LIKE_POST": self._describe_like_post,
|
||||
"DISLIKE_POST": self._describe_dislike_post,
|
||||
"REPOST": self._describe_repost,
|
||||
"QUOTE_POST": self._describe_quote_post,
|
||||
"FOLLOW": self._describe_follow,
|
||||
"CREATE_COMMENT": self._describe_create_comment,
|
||||
"LIKE_COMMENT": self._describe_like_comment,
|
||||
"DISLIKE_COMMENT": self._describe_dislike_comment,
|
||||
"SEARCH_POSTS": self._describe_search,
|
||||
"SEARCH_USER": self._describe_search_user,
|
||||
"MUTE": self._describe_mute,
|
||||
}
|
||||
|
||||
describe_func = action_descriptions.get(self.action_type, self._describe_generic)
|
||||
description = describe_func()
|
||||
|
||||
# 直接返回 "agent名称: 活动描述" 格式,不添加模拟前缀
|
||||
return f"{self.agent_name}: {description}"
|
||||
|
||||
def _describe_create_post(self) -> str:
|
||||
content = self.action_args.get("content", "")
|
||||
if content:
|
||||
return f"发布了一条帖子:「{content}」"
|
||||
return "发布了一条帖子"
|
||||
|
||||
def _describe_like_post(self) -> str:
|
||||
"""点赞帖子 - 包含帖子原文和作者信息"""
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if post_content and post_author:
|
||||
return f"点赞了{post_author}的帖子:「{post_content}」"
|
||||
elif post_content:
|
||||
return f"点赞了一条帖子:「{post_content}」"
|
||||
elif post_author:
|
||||
return f"点赞了{post_author}的一条帖子"
|
||||
return "点赞了一条帖子"
|
||||
|
||||
def _describe_dislike_post(self) -> str:
|
||||
"""踩帖子 - 包含帖子原文和作者信息"""
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if post_content and post_author:
|
||||
return f"踩了{post_author}的帖子:「{post_content}」"
|
||||
elif post_content:
|
||||
return f"踩了一条帖子:「{post_content}」"
|
||||
elif post_author:
|
||||
return f"踩了{post_author}的一条帖子"
|
||||
return "踩了一条帖子"
|
||||
|
||||
def _describe_repost(self) -> str:
|
||||
"""转发帖子 - 包含原帖内容和作者信息"""
|
||||
original_content = self.action_args.get("original_content", "")
|
||||
original_author = self.action_args.get("original_author_name", "")
|
||||
|
||||
if original_content and original_author:
|
||||
return f"转发了{original_author}的帖子:「{original_content}」"
|
||||
elif original_content:
|
||||
return f"转发了一条帖子:「{original_content}」"
|
||||
elif original_author:
|
||||
return f"转发了{original_author}的一条帖子"
|
||||
return "转发了一条帖子"
|
||||
|
||||
def _describe_quote_post(self) -> str:
|
||||
"""引用帖子 - 包含原帖内容、作者信息和引用评论"""
|
||||
original_content = self.action_args.get("original_content", "")
|
||||
original_author = self.action_args.get("original_author_name", "")
|
||||
quote_content = self.action_args.get("quote_content", "") or self.action_args.get("content", "")
|
||||
|
||||
base = ""
|
||||
if original_content and original_author:
|
||||
base = f"引用了{original_author}的帖子「{original_content}」"
|
||||
elif original_content:
|
||||
base = f"引用了一条帖子「{original_content}」"
|
||||
elif original_author:
|
||||
base = f"引用了{original_author}的一条帖子"
|
||||
else:
|
||||
base = "引用了一条帖子"
|
||||
|
||||
if quote_content:
|
||||
base += f",并评论道:「{quote_content}」"
|
||||
return base
|
||||
|
||||
def _describe_follow(self) -> str:
|
||||
"""关注用户 - 包含被关注用户的名称"""
|
||||
target_user_name = self.action_args.get("target_user_name", "")
|
||||
|
||||
if target_user_name:
|
||||
return f"关注了用户「{target_user_name}」"
|
||||
return "关注了一个用户"
|
||||
|
||||
def _describe_create_comment(self) -> str:
|
||||
"""发表评论 - 包含评论内容和所评论的帖子信息"""
|
||||
content = self.action_args.get("content", "")
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if content:
|
||||
if post_content and post_author:
|
||||
return f"在{post_author}的帖子「{post_content}」下评论道:「{content}」"
|
||||
elif post_content:
|
||||
return f"在帖子「{post_content}」下评论道:「{content}」"
|
||||
elif post_author:
|
||||
return f"在{post_author}的帖子下评论道:「{content}」"
|
||||
return f"评论道:「{content}」"
|
||||
return "发表了评论"
|
||||
|
||||
def _describe_like_comment(self) -> str:
|
||||
"""点赞评论 - 包含评论内容和作者信息"""
|
||||
comment_content = self.action_args.get("comment_content", "")
|
||||
comment_author = self.action_args.get("comment_author_name", "")
|
||||
|
||||
if comment_content and comment_author:
|
||||
return f"点赞了{comment_author}的评论:「{comment_content}」"
|
||||
elif comment_content:
|
||||
return f"点赞了一条评论:「{comment_content}」"
|
||||
elif comment_author:
|
||||
return f"点赞了{comment_author}的一条评论"
|
||||
return "点赞了一条评论"
|
||||
|
||||
def _describe_dislike_comment(self) -> str:
|
||||
"""踩评论 - 包含评论内容和作者信息"""
|
||||
comment_content = self.action_args.get("comment_content", "")
|
||||
comment_author = self.action_args.get("comment_author_name", "")
|
||||
|
||||
if comment_content and comment_author:
|
||||
return f"踩了{comment_author}的评论:「{comment_content}」"
|
||||
elif comment_content:
|
||||
return f"踩了一条评论:「{comment_content}」"
|
||||
elif comment_author:
|
||||
return f"踩了{comment_author}的一条评论"
|
||||
return "踩了一条评论"
|
||||
|
||||
def _describe_search(self) -> str:
|
||||
"""搜索帖子 - 包含搜索关键词"""
|
||||
query = self.action_args.get("query", "") or self.action_args.get("keyword", "")
|
||||
return f"搜索了「{query}」" if query else "进行了搜索"
|
||||
|
||||
def _describe_search_user(self) -> str:
|
||||
"""搜索用户 - 包含搜索关键词"""
|
||||
query = self.action_args.get("query", "") or self.action_args.get("username", "")
|
||||
return f"搜索了用户「{query}」" if query else "搜索了用户"
|
||||
|
||||
def _describe_mute(self) -> str:
|
||||
"""屏蔽用户 - 包含被屏蔽用户的名称"""
|
||||
target_user_name = self.action_args.get("target_user_name", "")
|
||||
|
||||
if target_user_name:
|
||||
return f"屏蔽了用户「{target_user_name}」"
|
||||
return "屏蔽了一个用户"
|
||||
|
||||
def _describe_generic(self) -> str:
|
||||
# 对于未知的动作类型,生成通用描述
|
||||
return f"执行了{self.action_type}操作"
|
||||
242
backend/app/services/memory_entity_reader.py
Normal file
242
backend/app/services/memory_entity_reader.py
Normal file
@@ -0,0 +1,242 @@
|
||||
"""Local entity-reader compatibility adapter for the legacy Zep consumer shape."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional
|
||||
|
||||
from .memory_repository import SqlAlchemyMemoryRepository
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalEntityNode:
|
||||
uuid: str
|
||||
name: str
|
||||
labels: list[str]
|
||||
summary: str
|
||||
attributes: dict[str, Any]
|
||||
related_edges: list[dict[str, Any]] = field(default_factory=list)
|
||||
related_nodes: list[dict[str, Any]] = field(default_factory=list)
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"uuid": self.uuid,
|
||||
"name": self.name,
|
||||
"labels": self.labels,
|
||||
"summary": self.summary,
|
||||
"attributes": self.attributes,
|
||||
"related_edges": self.related_edges,
|
||||
"related_nodes": self.related_nodes,
|
||||
}
|
||||
|
||||
def get_entity_type(self) -> Optional[str]:
|
||||
return next((label for label in self.labels if label not in {"Entity", "Node"}), None)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalFilteredEntities:
|
||||
entities: list[LocalEntityNode]
|
||||
entity_types: set[str]
|
||||
total_count: int
|
||||
filtered_count: int
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"entities": [entity.to_dict() for entity in self.entities],
|
||||
"entity_types": sorted(self.entity_types),
|
||||
"total_count": self.total_count,
|
||||
"filtered_count": self.filtered_count,
|
||||
}
|
||||
|
||||
|
||||
class LocalEntityReader:
|
||||
def __init__(
|
||||
self,
|
||||
session_or_repository,
|
||||
*,
|
||||
organization_id: str | None = None,
|
||||
graph_id: str | None = None,
|
||||
owns_session: bool = False,
|
||||
):
|
||||
self.owns_session = False
|
||||
if isinstance(session_or_repository, SqlAlchemyMemoryRepository):
|
||||
self.repository = session_or_repository
|
||||
else:
|
||||
if organization_id is None or graph_id is None:
|
||||
raise ValueError("memory_reader_scope_required")
|
||||
self.repository = SqlAlchemyMemoryRepository(
|
||||
session_or_repository,
|
||||
organization_id=organization_id,
|
||||
graph_id=graph_id,
|
||||
)
|
||||
self.owns_session = owns_session
|
||||
|
||||
def _validate_graph_id(self, graph_id: str | None = None) -> None:
|
||||
if graph_id is not None and graph_id != self.repository.graph_id:
|
||||
raise ValueError("memory_graph_scope_conflict")
|
||||
|
||||
def close(self) -> None:
|
||||
if self.owns_session:
|
||||
self.repository.session.close()
|
||||
self.owns_session = False
|
||||
|
||||
@staticmethod
|
||||
def _node_dict(node) -> dict[str, Any]:
|
||||
return {
|
||||
"uuid": node.id,
|
||||
"name": node.canonical_name,
|
||||
"labels": list(node.labels or []),
|
||||
"summary": node.summary or "",
|
||||
"attributes": dict(node.attributes or {}),
|
||||
}
|
||||
|
||||
def _entity(self, node, *, enrich_with_edges: bool = True) -> LocalEntityNode:
|
||||
related_edges: list[dict[str, Any]] = []
|
||||
related_nodes: list[dict[str, Any]] = []
|
||||
if enrich_with_edges:
|
||||
all_nodes = {item.id: item for item in self.repository.list_nodes()}
|
||||
for edge in self.repository.get_node_edges(node.id):
|
||||
source = all_nodes.get(edge.source_node_id)
|
||||
target = all_nodes.get(edge.target_node_id)
|
||||
if edge.source_node_id == node.id:
|
||||
related_edges.append(
|
||||
{
|
||||
"direction": "outgoing",
|
||||
"edge_name": edge.relation,
|
||||
"fact": edge.fact,
|
||||
"target_node_uuid": edge.target_node_id,
|
||||
}
|
||||
)
|
||||
related_node = target
|
||||
else:
|
||||
related_edges.append(
|
||||
{
|
||||
"direction": "incoming",
|
||||
"edge_name": edge.relation,
|
||||
"fact": edge.fact,
|
||||
"source_node_uuid": edge.source_node_id,
|
||||
}
|
||||
)
|
||||
related_node = source
|
||||
if related_node is not None:
|
||||
related_nodes.append(
|
||||
{
|
||||
"uuid": related_node.id,
|
||||
"name": related_node.canonical_name,
|
||||
"labels": list(related_node.labels or []),
|
||||
"summary": related_node.summary or "",
|
||||
}
|
||||
)
|
||||
return LocalEntityNode(
|
||||
uuid=node.id,
|
||||
name=node.canonical_name,
|
||||
labels=list(node.labels or []),
|
||||
summary=node.summary or "",
|
||||
attributes=dict(node.attributes or {}),
|
||||
related_edges=related_edges,
|
||||
related_nodes=related_nodes,
|
||||
)
|
||||
|
||||
def get_all_nodes(self, graph_id: str | None = None) -> list[dict[str, Any]]:
|
||||
self._validate_graph_id(graph_id)
|
||||
return [self._node_dict(node) for node in self.repository.list_nodes()]
|
||||
|
||||
def get_all_edges(self, graph_id: str | None = None) -> list[dict[str, Any]]:
|
||||
self._validate_graph_id(graph_id)
|
||||
return [
|
||||
{
|
||||
"uuid": edge.id,
|
||||
"name": edge.relation,
|
||||
"fact": edge.fact,
|
||||
"source_node_uuid": edge.source_node_id,
|
||||
"target_node_uuid": edge.target_node_id,
|
||||
"attributes": dict(edge.attributes or {}),
|
||||
}
|
||||
for edge in self.repository.list_edges()
|
||||
]
|
||||
|
||||
def filter_defined_entities(
|
||||
self,
|
||||
graph_id: str | None = None,
|
||||
*,
|
||||
defined_entity_types: Optional[list[str]] = None,
|
||||
enrich_with_edges: bool = True,
|
||||
) -> LocalFilteredEntities:
|
||||
self._validate_graph_id(graph_id)
|
||||
nodes = self.repository.list_nodes()
|
||||
allowed = set(defined_entity_types or [])
|
||||
entity_types: set[str] = set()
|
||||
entities: list[LocalEntityNode] = []
|
||||
for node in nodes:
|
||||
custom_labels = [label for label in node.labels or [] if label not in {"Entity", "Node"}]
|
||||
if not custom_labels:
|
||||
continue
|
||||
if defined_entity_types:
|
||||
matching_labels = [label for label in custom_labels if label in allowed]
|
||||
if not matching_labels:
|
||||
continue
|
||||
entity_type = matching_labels[0]
|
||||
else:
|
||||
entity_type = custom_labels[0]
|
||||
entity_types.add(entity_type)
|
||||
entities.append(self._entity(node, enrich_with_edges=enrich_with_edges))
|
||||
return LocalFilteredEntities(
|
||||
entities=entities,
|
||||
entity_types=entity_types,
|
||||
total_count=len(nodes),
|
||||
filtered_count=len(entities),
|
||||
)
|
||||
|
||||
def get_entity_with_context(
|
||||
self,
|
||||
graph_id_or_entity_uuid: str | None = None,
|
||||
entity_uuid: str | None = None,
|
||||
*,
|
||||
graph_id: str | None = None,
|
||||
) -> Optional[LocalEntityNode]:
|
||||
requested_graph_id = graph_id
|
||||
if entity_uuid is not None and requested_graph_id is None:
|
||||
requested_graph_id = graph_id_or_entity_uuid
|
||||
self._validate_graph_id(requested_graph_id)
|
||||
node_id = entity_uuid or graph_id_or_entity_uuid
|
||||
if not node_id:
|
||||
raise ValueError("memory_entity_id_required")
|
||||
node = self.repository.get_node(node_id)
|
||||
return self._entity(node, enrich_with_edges=True) if node is not None else None
|
||||
|
||||
def get_entities_by_type(
|
||||
self,
|
||||
graph_id_or_entity_type: str | None = None,
|
||||
entity_type: str | None = None,
|
||||
*,
|
||||
graph_id: str | None = None,
|
||||
enrich_with_edges: bool = True,
|
||||
) -> list[LocalEntityNode]:
|
||||
requested_graph_id = graph_id
|
||||
if entity_type is not None and requested_graph_id is None:
|
||||
requested_graph_id = graph_id_or_entity_type
|
||||
self._validate_graph_id(requested_graph_id)
|
||||
selected_type = entity_type or graph_id_or_entity_type
|
||||
if not selected_type:
|
||||
raise ValueError("memory_entity_type_required")
|
||||
return [
|
||||
self._entity(node, enrich_with_edges=enrich_with_edges)
|
||||
for node in self.repository.list_nodes()
|
||||
if selected_type in (node.labels or [])
|
||||
]
|
||||
|
||||
|
||||
def make_local_entity_reader_factory(session_factory, *, organization_id: str):
|
||||
"""Create per-worker readers; each reader owns and closes its own session."""
|
||||
if not isinstance(organization_id, str) or not organization_id.strip():
|
||||
raise ValueError("memory_reader_organization_required")
|
||||
|
||||
def factory(graph_id: str) -> LocalEntityReader:
|
||||
return LocalEntityReader(
|
||||
session_factory(),
|
||||
organization_id=organization_id,
|
||||
graph_id=graph_id,
|
||||
owns_session=True,
|
||||
)
|
||||
|
||||
return factory
|
||||
143
backend/app/services/memory_extraction.py
Normal file
143
backend/app/services/memory_extraction.py
Normal file
@@ -0,0 +1,143 @@
|
||||
"""Strict LLM contract for extracting local graph memory."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
from typing import Any
|
||||
|
||||
from pydantic import BaseModel, ConfigDict, Field, ValidationError
|
||||
|
||||
|
||||
MAX_EPISODE_CHARS = 20_000
|
||||
MAX_CONTEXT_CHARS = 12_000
|
||||
|
||||
|
||||
class _StrictModel(BaseModel):
|
||||
model_config = ConfigDict(extra="forbid")
|
||||
|
||||
|
||||
class ExtractedEntity(_StrictModel):
|
||||
mention: str = Field(min_length=1, max_length=2_000)
|
||||
canonical_name: str = Field(min_length=1, max_length=512)
|
||||
labels: list[str] = Field(default_factory=list, max_length=32)
|
||||
aliases: list[str] = Field(default_factory=list, max_length=32)
|
||||
attributes: dict[str, Any] = Field(default_factory=dict)
|
||||
summary: str = Field(default="", max_length=4_000)
|
||||
confidence: float = Field(default=0.0, ge=0.0, le=1.0)
|
||||
|
||||
|
||||
class ExtractedEdge(_StrictModel):
|
||||
source_entity_ref: str = Field(min_length=1, max_length=512)
|
||||
target_entity_ref: str = Field(min_length=1, max_length=512)
|
||||
relation: str = Field(min_length=1, max_length=128)
|
||||
fact: str = Field(min_length=1, max_length=4_000)
|
||||
attributes: dict[str, Any] = Field(default_factory=dict)
|
||||
valid_at: str | None = Field(default=None, max_length=128)
|
||||
invalid_at: str | None = Field(default=None, max_length=128)
|
||||
expired_at: str | None = Field(default=None, max_length=128)
|
||||
confidence: float = Field(default=0.0, ge=0.0, le=1.0)
|
||||
evidence: list[str] = Field(default_factory=list, max_length=32)
|
||||
|
||||
|
||||
class MemoryExtractionResult(_StrictModel):
|
||||
entities: list[ExtractedEntity] = Field(default_factory=list, max_length=500)
|
||||
edges: list[ExtractedEdge] = Field(default_factory=list, max_length=1_000)
|
||||
episode_summary: str = Field(default="", max_length=8_000)
|
||||
unresolved_mentions: list[str] = Field(default_factory=list, max_length=200)
|
||||
|
||||
|
||||
def parse_extraction_response(raw: str | dict[str, Any]) -> MemoryExtractionResult:
|
||||
"""Parse and validate an LLM response without accepting extra fields."""
|
||||
if isinstance(raw, str):
|
||||
text = raw.strip()
|
||||
if text.startswith("```"):
|
||||
lines = text.splitlines()
|
||||
if lines and lines[0].strip().startswith("```"):
|
||||
lines = lines[1:]
|
||||
if lines and lines[-1].strip() == "```":
|
||||
lines = lines[:-1]
|
||||
text = "\n".join(lines).strip()
|
||||
try:
|
||||
raw = json.loads(text)
|
||||
except json.JSONDecodeError as exc:
|
||||
raise ValueError("invalid_memory_json") from exc
|
||||
if not isinstance(raw, dict):
|
||||
raise ValueError("invalid_memory_payload")
|
||||
return MemoryExtractionResult.model_validate(raw)
|
||||
|
||||
|
||||
def _language_instruction(language: str) -> str:
|
||||
if language == "en":
|
||||
return "IMPORTANT: Write summaries and facts in English. Return JSON only."
|
||||
return "IMPORTANT: Write summaries and facts in Thai. Return JSON only."
|
||||
|
||||
|
||||
def build_extraction_prompt(
|
||||
*,
|
||||
language: str,
|
||||
ontology: dict[str, Any],
|
||||
episode_text: str,
|
||||
context: str = "",
|
||||
) -> str:
|
||||
if not isinstance(episode_text, str) or len(episode_text) > MAX_EPISODE_CHARS:
|
||||
raise ValueError("episode_too_large")
|
||||
if not isinstance(context, str) or len(context) > MAX_CONTEXT_CHARS:
|
||||
raise ValueError("context_too_large")
|
||||
if not isinstance(ontology, dict):
|
||||
raise ValueError("invalid_ontology")
|
||||
|
||||
ontology_json = json.dumps(ontology, ensure_ascii=False, sort_keys=True)
|
||||
context_block = context if context else "(none)"
|
||||
return f"""{_language_instruction(language)}
|
||||
|
||||
You extract evidence-grounded graph memory from one episode.
|
||||
Never invent facts. Use only labels and relations allowed by the ontology.
|
||||
Use stable entity_refs inside this response; never guess database IDs.
|
||||
Preserve temporal fields as null when the evidence does not support a date.
|
||||
Keep confidence between 0 and 1 and keep evidence references when available.
|
||||
|
||||
Return JSON only with this shape:
|
||||
{{
|
||||
"entities": [{{
|
||||
"mention": "text span",
|
||||
"canonical_name": "stable name",
|
||||
"labels": ["Person"],
|
||||
"aliases": [],
|
||||
"attributes": {{}},
|
||||
"summary": "short evidence-grounded summary",
|
||||
"confidence": 0.0
|
||||
}}],
|
||||
"edges": [{{
|
||||
"source_entity_ref": "entity-ref",
|
||||
"target_entity_ref": "entity-ref",
|
||||
"relation": "RELATION_NAME",
|
||||
"fact": "evidence-grounded fact",
|
||||
"attributes": {{}},
|
||||
"valid_at": null,
|
||||
"invalid_at": null,
|
||||
"expired_at": null,
|
||||
"confidence": 0.0,
|
||||
"evidence": ["episode reference or span"]
|
||||
}}],
|
||||
"episode_summary": "short summary",
|
||||
"unresolved_mentions": []
|
||||
}}
|
||||
|
||||
Ontology:
|
||||
{ontology_json}
|
||||
|
||||
Additional context:
|
||||
{context_block}
|
||||
|
||||
Episode:
|
||||
{episode_text}
|
||||
"""
|
||||
|
||||
|
||||
__all__ = [
|
||||
"ExtractedEdge",
|
||||
"ExtractedEntity",
|
||||
"MemoryExtractionResult",
|
||||
"build_extraction_prompt",
|
||||
"parse_extraction_response",
|
||||
]
|
||||
310
backend/app/services/memory_repository.py
Normal file
310
backend/app/services/memory_repository.py
Normal file
@@ -0,0 +1,310 @@
|
||||
"""Tenant-scoped SQLAlchemy repository for local graph memory."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import math
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any
|
||||
|
||||
from sqlalchemy import or_, select
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.memory import MemoryEdge, MemoryEpisode, MemoryGraph, MemoryNode
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MemorySearchResult:
|
||||
nodes: list[MemoryNode] = field(default_factory=list)
|
||||
edges: list[MemoryEdge] = field(default_factory=list)
|
||||
query: str = ""
|
||||
|
||||
@property
|
||||
def total_count(self) -> int:
|
||||
return len(self.nodes) + len(self.edges)
|
||||
|
||||
|
||||
class SqlAlchemyMemoryRepository:
|
||||
"""All reads and writes are constrained by organization_id + graph_id."""
|
||||
|
||||
def __init__(self, session: Session, *, organization_id: str, graph_id: str):
|
||||
if not isinstance(organization_id, str) or not organization_id.strip():
|
||||
raise ValueError("invalid_organization_id")
|
||||
if not isinstance(graph_id, str) or not graph_id.strip():
|
||||
raise ValueError("invalid_graph_id")
|
||||
self.session = session
|
||||
self.organization_id = organization_id
|
||||
self.graph_id = graph_id
|
||||
|
||||
def _graph(self) -> MemoryGraph | None:
|
||||
return self.session.scalar(
|
||||
select(MemoryGraph).where(
|
||||
MemoryGraph.id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
)
|
||||
)
|
||||
|
||||
def _require_graph(self) -> MemoryGraph:
|
||||
graph = self._graph()
|
||||
if graph is None:
|
||||
raise ValueError("memory_graph_not_found")
|
||||
return graph
|
||||
|
||||
def get_graph(self) -> MemoryGraph:
|
||||
"""Return the graph only when it belongs to this organization scope."""
|
||||
return self._require_graph()
|
||||
|
||||
def update_graph_ontology(self, ontology: dict[str, Any]) -> MemoryGraph:
|
||||
if not isinstance(ontology, dict):
|
||||
raise ValueError("invalid_ontology")
|
||||
graph = self._require_graph()
|
||||
graph.ontology = dict(ontology)
|
||||
graph.version = int(graph.version or 0) + 1
|
||||
self.session.flush()
|
||||
return graph
|
||||
|
||||
def get_episode(self, *, source_type: str, source_ref: str) -> MemoryEpisode | None:
|
||||
self._require_graph()
|
||||
return self.session.scalar(
|
||||
select(MemoryEpisode).where(
|
||||
MemoryEpisode.graph_id == self.graph_id,
|
||||
MemoryEpisode.source_type == source_type,
|
||||
MemoryEpisode.source_ref == source_ref,
|
||||
)
|
||||
)
|
||||
|
||||
@staticmethod
|
||||
def _normalize_name(name: str) -> str:
|
||||
if not isinstance(name, str) or not name.strip():
|
||||
raise ValueError("invalid_memory_name")
|
||||
return " ".join(name.casefold().split())
|
||||
|
||||
@staticmethod
|
||||
def _confidence(value: float) -> float:
|
||||
if isinstance(value, bool) or not isinstance(value, (int, float)) or not math.isfinite(value):
|
||||
raise ValueError("invalid_confidence")
|
||||
return min(max(float(value), 0.0), 1.0)
|
||||
|
||||
def create_graph(self, *, project_id: str, ontology: dict[str, Any] | None = None) -> MemoryGraph:
|
||||
if not isinstance(project_id, str) or not project_id.strip():
|
||||
raise ValueError("invalid_project_id")
|
||||
existing = self.session.get(MemoryGraph, self.graph_id)
|
||||
if existing is not None:
|
||||
if existing.organization_id != self.organization_id:
|
||||
raise ValueError("memory_graph_scope_conflict")
|
||||
return existing
|
||||
graph = MemoryGraph(
|
||||
id=self.graph_id,
|
||||
organization_id=self.organization_id,
|
||||
project_id=project_id,
|
||||
ontology=ontology or {},
|
||||
)
|
||||
self.session.add(graph)
|
||||
self.session.flush()
|
||||
return graph
|
||||
|
||||
def add_episode(
|
||||
self,
|
||||
*,
|
||||
source_type: str,
|
||||
source_ref: str,
|
||||
normalized_text: str,
|
||||
summary: str = "",
|
||||
extractor_version: str = "v1",
|
||||
) -> MemoryEpisode:
|
||||
self._require_graph()
|
||||
if not all(isinstance(value, str) and value.strip() for value in (source_type, source_ref, normalized_text)):
|
||||
raise ValueError("invalid_memory_episode")
|
||||
episode = self.session.scalar(
|
||||
select(MemoryEpisode).where(
|
||||
MemoryEpisode.graph_id == self.graph_id,
|
||||
MemoryEpisode.source_type == source_type,
|
||||
MemoryEpisode.source_ref == source_ref,
|
||||
)
|
||||
)
|
||||
if episode is None:
|
||||
episode = MemoryEpisode(
|
||||
graph_id=self.graph_id,
|
||||
source_type=source_type,
|
||||
source_ref=source_ref,
|
||||
normalized_text=normalized_text,
|
||||
summary=summary or "",
|
||||
extractor_version=extractor_version or "v1",
|
||||
)
|
||||
self.session.add(episode)
|
||||
else:
|
||||
episode.normalized_text = normalized_text
|
||||
episode.summary = summary or ""
|
||||
episode.extractor_version = extractor_version or "v1"
|
||||
self.session.flush()
|
||||
return episode
|
||||
|
||||
def upsert_node(
|
||||
self,
|
||||
*,
|
||||
canonical_name: str,
|
||||
labels: list[str] | None = None,
|
||||
aliases: list[str] | None = None,
|
||||
attributes: dict[str, Any] | None = None,
|
||||
summary: str = "",
|
||||
confidence: float = 0.0,
|
||||
) -> MemoryNode:
|
||||
self._require_graph()
|
||||
normalized_name = self._normalize_name(canonical_name)
|
||||
node = self.session.scalar(
|
||||
select(MemoryNode).where(
|
||||
MemoryNode.graph_id == self.graph_id,
|
||||
MemoryNode.normalized_name == normalized_name,
|
||||
)
|
||||
)
|
||||
if node is None:
|
||||
node = MemoryNode(graph_id=self.graph_id, canonical_name=canonical_name.strip(), normalized_name=normalized_name)
|
||||
self.session.add(node)
|
||||
node.canonical_name = canonical_name.strip()
|
||||
node.labels = list(labels or [])
|
||||
node.aliases = list(aliases or [])
|
||||
node.attributes = dict(attributes or {})
|
||||
node.summary = summary or ""
|
||||
node.confidence = self._confidence(confidence)
|
||||
self.session.flush()
|
||||
return node
|
||||
|
||||
def upsert_edge(
|
||||
self,
|
||||
*,
|
||||
source_node_id: str,
|
||||
target_node_id: str,
|
||||
relation: str,
|
||||
fact: str,
|
||||
attributes: dict[str, Any] | None = None,
|
||||
confidence: float = 0.0,
|
||||
valid_at=None,
|
||||
invalid_at=None,
|
||||
expired_at=None,
|
||||
) -> MemoryEdge:
|
||||
self._require_graph()
|
||||
if not all(isinstance(value, str) and value.strip() for value in (source_node_id, target_node_id, relation, fact)):
|
||||
raise ValueError("invalid_memory_edge")
|
||||
nodes = list(
|
||||
self.session.scalars(
|
||||
select(MemoryNode).where(
|
||||
MemoryNode.graph_id == self.graph_id,
|
||||
MemoryNode.id.in_([source_node_id, target_node_id]),
|
||||
)
|
||||
)
|
||||
)
|
||||
if {node.id for node in nodes} != {source_node_id, target_node_id}:
|
||||
raise ValueError("memory_edge_node_scope_conflict")
|
||||
edge = self.session.scalar(
|
||||
select(MemoryEdge).where(
|
||||
MemoryEdge.graph_id == self.graph_id,
|
||||
MemoryEdge.source_node_id == source_node_id,
|
||||
MemoryEdge.target_node_id == target_node_id,
|
||||
MemoryEdge.relation == relation.strip(),
|
||||
MemoryEdge.fact == fact.strip(),
|
||||
)
|
||||
)
|
||||
if edge is None:
|
||||
edge = MemoryEdge(
|
||||
graph_id=self.graph_id,
|
||||
source_node_id=source_node_id,
|
||||
target_node_id=target_node_id,
|
||||
relation=relation.strip(),
|
||||
fact=fact.strip(),
|
||||
)
|
||||
self.session.add(edge)
|
||||
edge.attributes = dict(attributes or {})
|
||||
edge.confidence = self._confidence(confidence)
|
||||
edge.valid_at = valid_at
|
||||
edge.invalid_at = invalid_at
|
||||
edge.expired_at = expired_at
|
||||
self.session.flush()
|
||||
return edge
|
||||
|
||||
def get_node(self, node_id: str) -> MemoryNode | None:
|
||||
return self.session.scalar(
|
||||
select(MemoryNode)
|
||||
.join(MemoryGraph, MemoryGraph.id == MemoryNode.graph_id)
|
||||
.where(
|
||||
MemoryNode.id == node_id,
|
||||
MemoryNode.graph_id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
)
|
||||
)
|
||||
|
||||
def list_nodes(self, *, limit: int = 1_000) -> list[MemoryNode]:
|
||||
self._require_graph()
|
||||
safe_limit = min(max(int(limit), 1), 10_000)
|
||||
return list(
|
||||
self.session.scalars(
|
||||
select(MemoryNode)
|
||||
.join(MemoryGraph, MemoryGraph.id == MemoryNode.graph_id)
|
||||
.where(
|
||||
MemoryNode.graph_id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
)
|
||||
.order_by(MemoryNode.normalized_name.asc(), MemoryNode.id.asc())
|
||||
.limit(safe_limit)
|
||||
)
|
||||
)
|
||||
|
||||
def list_edges(self, *, limit: int = 2_000) -> list[MemoryEdge]:
|
||||
self._require_graph()
|
||||
safe_limit = min(max(int(limit), 1), 20_000)
|
||||
return list(
|
||||
self.session.scalars(
|
||||
select(MemoryEdge)
|
||||
.join(MemoryGraph, MemoryGraph.id == MemoryEdge.graph_id)
|
||||
.where(
|
||||
MemoryEdge.graph_id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
)
|
||||
.order_by(MemoryEdge.id.asc())
|
||||
.limit(safe_limit)
|
||||
)
|
||||
)
|
||||
|
||||
def get_node_edges(self, node_id: str, *, limit: int = 500) -> list[MemoryEdge]:
|
||||
self._require_graph()
|
||||
safe_limit = min(max(int(limit), 1), 5_000)
|
||||
return list(
|
||||
self.session.scalars(
|
||||
select(MemoryEdge)
|
||||
.join(MemoryGraph, MemoryGraph.id == MemoryEdge.graph_id)
|
||||
.where(
|
||||
MemoryEdge.graph_id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
(MemoryEdge.source_node_id == node_id) | (MemoryEdge.target_node_id == node_id),
|
||||
)
|
||||
.order_by(MemoryEdge.id.asc())
|
||||
.limit(safe_limit)
|
||||
)
|
||||
)
|
||||
|
||||
def get_node_by_id(self, node_id: str) -> MemoryNode | None:
|
||||
return self.get_node(node_id)
|
||||
|
||||
def search(self, query: str, *, limit: int = 50) -> MemorySearchResult:
|
||||
if not isinstance(query, str):
|
||||
raise ValueError("invalid_memory_query")
|
||||
normalized_query = " ".join(query.casefold().split())
|
||||
if not normalized_query:
|
||||
return MemorySearchResult(query=query)
|
||||
safe_limit = min(max(int(limit), 1), 100)
|
||||
pattern = f"%{normalized_query}%"
|
||||
nodes = list(
|
||||
self.session.scalars(
|
||||
select(MemoryNode)
|
||||
.join(MemoryGraph, MemoryGraph.id == MemoryNode.graph_id)
|
||||
.where(
|
||||
MemoryNode.graph_id == self.graph_id,
|
||||
MemoryGraph.organization_id == self.organization_id,
|
||||
or_(
|
||||
MemoryNode.normalized_name.ilike(pattern),
|
||||
MemoryNode.summary.ilike(pattern),
|
||||
),
|
||||
)
|
||||
.order_by(MemoryNode.normalized_name.asc(), MemoryNode.id.asc())
|
||||
.limit(safe_limit)
|
||||
)
|
||||
)
|
||||
return MemorySearchResult(nodes=nodes, query=query)
|
||||
122
backend/app/services/memory_service.py
Normal file
122
backend/app/services/memory_service.py
Normal file
@@ -0,0 +1,122 @@
|
||||
"""LLM extraction orchestration without giving the model storage authority."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime, timezone
|
||||
from typing import Any, Protocol
|
||||
|
||||
from .memory_extraction import MemoryExtractionResult, build_extraction_prompt, parse_extraction_response
|
||||
from .memory_repository import SqlAlchemyMemoryRepository
|
||||
|
||||
|
||||
class JsonLLMClient(Protocol):
|
||||
def chat_json(self, messages: list[dict[str, str]], temperature: float = 0.3, max_tokens: int = 4096) -> dict[str, Any]:
|
||||
...
|
||||
|
||||
|
||||
@dataclass(frozen=True)
|
||||
class MemoryIngestResult:
|
||||
entity_count: int
|
||||
edge_count: int
|
||||
unresolved_edge_refs: list[str] = field(default_factory=list)
|
||||
|
||||
|
||||
class MemoryExtractionService:
|
||||
def __init__(self, client: JsonLLMClient):
|
||||
self.client = client
|
||||
|
||||
def extract(
|
||||
self,
|
||||
*,
|
||||
language: str,
|
||||
ontology: dict[str, Any],
|
||||
episode_text: str,
|
||||
context: str = "",
|
||||
) -> MemoryExtractionResult:
|
||||
prompt = build_extraction_prompt(
|
||||
language=language,
|
||||
ontology=ontology,
|
||||
episode_text=episode_text,
|
||||
context=context,
|
||||
)
|
||||
raw = self.client.chat_json(
|
||||
[{"role": "system", "content": prompt}],
|
||||
temperature=0.2,
|
||||
max_tokens=8192,
|
||||
)
|
||||
return parse_extraction_response(raw)
|
||||
|
||||
@staticmethod
|
||||
def _ref(value: str) -> str:
|
||||
return " ".join(value.casefold().split())
|
||||
|
||||
@staticmethod
|
||||
def _timestamp(value: str | None):
|
||||
if value is None:
|
||||
return None
|
||||
if not isinstance(value, str) or not value.strip():
|
||||
raise ValueError("invalid_memory_timestamp")
|
||||
try:
|
||||
parsed = datetime.fromisoformat(value.strip().replace("Z", "+00:00"))
|
||||
except ValueError as exc:
|
||||
raise ValueError("invalid_memory_timestamp") from exc
|
||||
return parsed if parsed.tzinfo is not None else parsed.replace(tzinfo=timezone.utc)
|
||||
|
||||
def persist(
|
||||
self,
|
||||
repository: SqlAlchemyMemoryRepository,
|
||||
result: MemoryExtractionResult,
|
||||
*,
|
||||
source_type: str,
|
||||
source_ref: str,
|
||||
episode_text: str,
|
||||
) -> MemoryIngestResult:
|
||||
repository.add_episode(
|
||||
source_type=source_type,
|
||||
source_ref=source_ref,
|
||||
normalized_text=episode_text,
|
||||
summary=result.episode_summary,
|
||||
)
|
||||
|
||||
entity_refs: dict[str, str] = {}
|
||||
entity_count = 0
|
||||
for entity in result.entities:
|
||||
node = repository.upsert_node(
|
||||
canonical_name=entity.canonical_name,
|
||||
labels=entity.labels,
|
||||
aliases=entity.aliases,
|
||||
attributes=entity.attributes,
|
||||
summary=entity.summary,
|
||||
confidence=entity.confidence,
|
||||
)
|
||||
entity_count += 1
|
||||
for reference in [entity.mention, entity.canonical_name, *entity.aliases]:
|
||||
entity_refs[self._ref(reference)] = node.id
|
||||
|
||||
unresolved: list[str] = []
|
||||
edge_count = 0
|
||||
for edge in result.edges:
|
||||
source_id = entity_refs.get(self._ref(edge.source_entity_ref))
|
||||
target_id = entity_refs.get(self._ref(edge.target_entity_ref))
|
||||
if source_id is None or target_id is None:
|
||||
unresolved.append(f"{edge.source_entity_ref}->{edge.target_entity_ref}")
|
||||
continue
|
||||
repository.upsert_edge(
|
||||
source_node_id=source_id,
|
||||
target_node_id=target_id,
|
||||
relation=edge.relation,
|
||||
fact=edge.fact,
|
||||
attributes=edge.attributes,
|
||||
confidence=edge.confidence,
|
||||
valid_at=self._timestamp(edge.valid_at),
|
||||
invalid_at=self._timestamp(edge.invalid_at),
|
||||
expired_at=self._timestamp(edge.expired_at),
|
||||
)
|
||||
edge_count += 1
|
||||
|
||||
return MemoryIngestResult(
|
||||
entity_count=entity_count,
|
||||
edge_count=edge_count,
|
||||
unresolved_edge_refs=unresolved,
|
||||
)
|
||||
506
backend/app/services/memory_tools.py
Normal file
506
backend/app/services/memory_tools.py
Normal file
@@ -0,0 +1,506 @@
|
||||
"""Local compatibility tools for the legacy ZepToolsService result contract."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Any, Optional
|
||||
|
||||
from .memory_repository import SqlAlchemyMemoryRepository
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalSearchResult:
|
||||
facts: list[str]
|
||||
edges: list[dict[str, Any]]
|
||||
nodes: list[dict[str, Any]]
|
||||
query: str
|
||||
total_count: int
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"facts": self.facts,
|
||||
"edges": self.edges,
|
||||
"nodes": self.nodes,
|
||||
"query": self.query,
|
||||
"total_count": self.total_count,
|
||||
}
|
||||
|
||||
def to_text(self) -> str:
|
||||
parts = [f"Search query: {self.query}", f"Found {self.total_count} relevant items"]
|
||||
if self.facts:
|
||||
parts.append("\n### Relevant facts:")
|
||||
parts.extend(f"{index}. {fact}" for index, fact in enumerate(self.facts, 1))
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalNodeInfo:
|
||||
uuid: str
|
||||
name: str
|
||||
labels: list[str]
|
||||
summary: str
|
||||
attributes: dict[str, Any]
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"uuid": self.uuid,
|
||||
"name": self.name,
|
||||
"labels": self.labels,
|
||||
"summary": self.summary,
|
||||
"attributes": self.attributes,
|
||||
}
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalInsightForgeResult:
|
||||
query: str
|
||||
simulation_requirement: str
|
||||
sub_queries: list[str]
|
||||
semantic_facts: list[str] = field(default_factory=list)
|
||||
entity_insights: list[dict[str, Any]] = field(default_factory=list)
|
||||
relationship_chains: list[str] = field(default_factory=list)
|
||||
total_facts: int = 0
|
||||
total_entities: int = 0
|
||||
total_relationships: int = 0
|
||||
|
||||
def to_text(self) -> str:
|
||||
parts = [
|
||||
"## Local Memory Deep Analysis",
|
||||
f"Analysis question: {self.query}",
|
||||
f"Prediction scenario: {self.simulation_requirement}",
|
||||
f"\n### Statistics\n- Facts: {self.total_facts}\n- Entities: {self.total_entities}\n- Relationships: {self.total_relationships}",
|
||||
]
|
||||
if self.semantic_facts:
|
||||
parts.append("\n### Key facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.semantic_facts, 1)))
|
||||
if self.entity_insights:
|
||||
parts.append(
|
||||
"\n### Entities\n"
|
||||
+ "\n".join(
|
||||
f"- {item.get('name', 'Unknown')} ({item.get('type', 'Entity')}): {item.get('summary', '')}"
|
||||
for item in self.entity_insights
|
||||
)
|
||||
)
|
||||
if self.relationship_chains:
|
||||
parts.append("\n### Relationships\n" + "\n".join(f"- {chain}" for chain in self.relationship_chains))
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalPanoramaResult:
|
||||
query: str
|
||||
all_nodes: list[LocalNodeInfo] = field(default_factory=list)
|
||||
all_edges: list[LocalEdgeInfo] = field(default_factory=list)
|
||||
active_facts: list[str] = field(default_factory=list)
|
||||
historical_facts: list[str] = field(default_factory=list)
|
||||
|
||||
def to_text(self) -> str:
|
||||
parts = [
|
||||
"## Local Memory Panorama",
|
||||
f"Query: {self.query}",
|
||||
f"\n### Statistics\n- Total nodes: {len(self.all_nodes)}\n- Total edges: {len(self.all_edges)}\n- Active facts: {len(self.active_facts)}\n- Historical facts: {len(self.historical_facts)}",
|
||||
]
|
||||
if self.active_facts:
|
||||
parts.append("\n### Active facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.active_facts, 1)))
|
||||
if self.historical_facts:
|
||||
parts.append("\n### Historical facts\n" + "\n".join(f"{i}. {fact}" for i, fact in enumerate(self.historical_facts, 1)))
|
||||
if self.all_nodes:
|
||||
parts.append("\n### Entities\n" + "\n".join(f"- {node.name}: {node.summary}" for node in self.all_nodes))
|
||||
return "\n".join(parts)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalInterviewResult:
|
||||
interview_topic: str
|
||||
summary: str
|
||||
|
||||
def to_text(self) -> str:
|
||||
return (
|
||||
"## Local Memory Interview\n"
|
||||
f"Topic: {self.interview_topic}\n\n"
|
||||
f"{self.summary}"
|
||||
)
|
||||
|
||||
|
||||
@dataclass
|
||||
class LocalEdgeInfo:
|
||||
uuid: str
|
||||
name: str
|
||||
fact: str
|
||||
source_node_uuid: str
|
||||
target_node_uuid: str
|
||||
source_node_name: Optional[str] = None
|
||||
target_node_name: Optional[str] = None
|
||||
created_at: Optional[str] = None
|
||||
valid_at: Optional[str] = None
|
||||
invalid_at: Optional[str] = None
|
||||
expired_at: Optional[str] = None
|
||||
|
||||
def to_dict(self) -> dict[str, Any]:
|
||||
return {
|
||||
"uuid": self.uuid,
|
||||
"name": self.name,
|
||||
"fact": self.fact,
|
||||
"source_node_uuid": self.source_node_uuid,
|
||||
"target_node_uuid": self.target_node_uuid,
|
||||
"source_node_name": self.source_node_name,
|
||||
"target_node_name": self.target_node_name,
|
||||
"created_at": self.created_at,
|
||||
"valid_at": self.valid_at,
|
||||
"invalid_at": self.invalid_at,
|
||||
"expired_at": self.expired_at,
|
||||
}
|
||||
|
||||
@property
|
||||
def is_expired(self) -> bool:
|
||||
return self.expired_at is not None
|
||||
|
||||
@property
|
||||
def is_invalid(self) -> bool:
|
||||
return self.invalid_at is not None
|
||||
|
||||
|
||||
class LocalMemoryTools:
|
||||
def __init__(self, session_or_repository, *, organization_id: str | None = None, graph_id: str | None = None):
|
||||
if isinstance(session_or_repository, SqlAlchemyMemoryRepository):
|
||||
self.repository = session_or_repository
|
||||
else:
|
||||
if organization_id is None or graph_id is None:
|
||||
raise ValueError("memory_tools_scope_required")
|
||||
self.repository = SqlAlchemyMemoryRepository(
|
||||
session_or_repository,
|
||||
organization_id=organization_id,
|
||||
graph_id=graph_id,
|
||||
)
|
||||
|
||||
def _validate_graph_id(self, graph_id: str | None = None) -> None:
|
||||
if graph_id is not None and graph_id != self.repository.graph_id:
|
||||
raise ValueError("memory_graph_scope_conflict")
|
||||
|
||||
@staticmethod
|
||||
def _node_info(node) -> LocalNodeInfo:
|
||||
return LocalNodeInfo(
|
||||
uuid=node.id,
|
||||
name=node.canonical_name,
|
||||
labels=list(node.labels or []),
|
||||
summary=node.summary or "",
|
||||
attributes=dict(node.attributes or {}),
|
||||
)
|
||||
|
||||
def _edge_info(self, edge, nodes_by_id: dict[str, Any]) -> LocalEdgeInfo:
|
||||
source = nodes_by_id.get(edge.source_node_id)
|
||||
target = nodes_by_id.get(edge.target_node_id)
|
||||
return LocalEdgeInfo(
|
||||
uuid=edge.id,
|
||||
name=edge.relation,
|
||||
fact=edge.fact,
|
||||
source_node_uuid=edge.source_node_id,
|
||||
target_node_uuid=edge.target_node_id,
|
||||
source_node_name=source.canonical_name if source else None,
|
||||
target_node_name=target.canonical_name if target else None,
|
||||
created_at=edge.created_at.isoformat() if edge.created_at else None,
|
||||
valid_at=edge.valid_at.isoformat() if edge.valid_at else None,
|
||||
invalid_at=edge.invalid_at.isoformat() if edge.invalid_at else None,
|
||||
expired_at=edge.expired_at.isoformat() if edge.expired_at else None,
|
||||
)
|
||||
|
||||
def get_all_nodes(self, *, limit: int = 10_000) -> list[LocalNodeInfo]:
|
||||
return [self._node_info(node) for node in self.repository.list_nodes(limit=limit)]
|
||||
|
||||
def get_all_edges(self, *, limit: int = 20_000) -> list[LocalEdgeInfo]:
|
||||
nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
|
||||
return [self._edge_info(edge, nodes) for edge in self.repository.list_edges(limit=limit)]
|
||||
|
||||
@staticmethod
|
||||
def _score(query: str, *values: str) -> int:
|
||||
normalized_query = query.casefold().strip()
|
||||
haystack = " ".join(value or "" for value in values).casefold()
|
||||
if not normalized_query or not haystack:
|
||||
return 0
|
||||
if normalized_query in haystack:
|
||||
return 100
|
||||
return sum(10 for token in normalized_query.split() if token in haystack)
|
||||
|
||||
def search_graph(
|
||||
self,
|
||||
query: str,
|
||||
*,
|
||||
limit: int = 10,
|
||||
scope: str = "edges",
|
||||
graph_id: str | None = None,
|
||||
) -> LocalSearchResult:
|
||||
self._validate_graph_id(graph_id)
|
||||
if scope not in {"edges", "nodes", "both"}:
|
||||
raise ValueError("invalid_memory_search_scope")
|
||||
safe_limit = min(max(int(limit), 1), 100)
|
||||
facts: list[str] = []
|
||||
edges: list[dict[str, Any]] = []
|
||||
nodes: list[dict[str, Any]] = []
|
||||
all_nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
|
||||
|
||||
if scope in {"edges", "both"}:
|
||||
scored_edges = [
|
||||
(self._score(query, edge.fact, edge.relation), edge)
|
||||
for edge in self.repository.list_edges(limit=20_000)
|
||||
]
|
||||
for score, edge in sorted(scored_edges, key=lambda item: (-item[0], item[1].id))[:safe_limit]:
|
||||
if score <= 0:
|
||||
continue
|
||||
edge_info = self._edge_info(edge, all_nodes).to_dict()
|
||||
edges.append(edge_info)
|
||||
if edge.fact:
|
||||
facts.append(edge.fact)
|
||||
|
||||
if scope in {"nodes", "both"}:
|
||||
scored_nodes = [
|
||||
(self._score(query, node.canonical_name, node.summary), node)
|
||||
for node in all_nodes.values()
|
||||
]
|
||||
for score, node in sorted(scored_nodes, key=lambda item: (-item[0], item[1].id))[:safe_limit]:
|
||||
if score <= 0:
|
||||
continue
|
||||
node_info = self._node_info(node).to_dict()
|
||||
node_info.pop("attributes", None)
|
||||
nodes.append(node_info)
|
||||
if node.summary:
|
||||
facts.append(f"[{node.canonical_name}]: {node.summary}")
|
||||
|
||||
return LocalSearchResult(
|
||||
facts=facts,
|
||||
edges=edges,
|
||||
nodes=nodes,
|
||||
query=query,
|
||||
total_count=len(facts),
|
||||
)
|
||||
|
||||
def quick_search(
|
||||
self,
|
||||
query: str,
|
||||
*,
|
||||
limit: int = 10,
|
||||
graph_id: str | None = None,
|
||||
) -> LocalSearchResult:
|
||||
return self.search_graph(query, limit=limit, scope="edges", graph_id=graph_id)
|
||||
|
||||
def get_node_detail(self, node_uuid: str) -> LocalNodeInfo | None:
|
||||
node = self.repository.get_node(node_uuid)
|
||||
return self._node_info(node) if node is not None else None
|
||||
|
||||
def get_node_edges(self, node_uuid: str, *, limit: int = 500) -> list[LocalEdgeInfo]:
|
||||
nodes = {node.id: node for node in self.repository.list_nodes(limit=10_000)}
|
||||
return [self._edge_info(edge, nodes) for edge in self.repository.get_node_edges(node_uuid, limit=limit)]
|
||||
|
||||
def insight_forge(
|
||||
self,
|
||||
*,
|
||||
graph_id: str,
|
||||
query: str,
|
||||
simulation_requirement: str = "",
|
||||
report_context: str = "",
|
||||
) -> LocalInsightForgeResult:
|
||||
self._validate_graph_id(graph_id)
|
||||
result = self.search_graph(query, limit=100, scope="edges", graph_id=graph_id)
|
||||
facts = list(dict.fromkeys(result.facts))
|
||||
nodes_by_id = {
|
||||
node.id: node for node in self.repository.list_nodes(limit=10_000)
|
||||
}
|
||||
related_node_ids = list(
|
||||
dict.fromkeys(
|
||||
node_id
|
||||
for edge in result.edges
|
||||
for node_id in (edge.get("source_node_uuid"), edge.get("target_node_uuid"))
|
||||
if node_id
|
||||
)
|
||||
)
|
||||
entity_insights = []
|
||||
for node_id in related_node_ids:
|
||||
node = nodes_by_id.get(node_id)
|
||||
if node is None:
|
||||
continue
|
||||
entity_insights.append(
|
||||
{
|
||||
"uuid": node.id,
|
||||
"name": node.canonical_name,
|
||||
"type": next(
|
||||
(label for label in (node.labels or []) if label not in {"Entity", "Node"}),
|
||||
"Entity",
|
||||
),
|
||||
"summary": node.summary or "",
|
||||
"related_facts": [
|
||||
fact for fact in facts if node.canonical_name.casefold() in fact.casefold()
|
||||
],
|
||||
}
|
||||
)
|
||||
|
||||
relationship_chains = []
|
||||
seen_chains = set()
|
||||
for edge in result.edges:
|
||||
chain = (
|
||||
f"{edge.get('source_node_name') or edge.get('source_node_uuid')} "
|
||||
f"--[{edge.get('name', '')}]--> "
|
||||
f"{edge.get('target_node_name') or edge.get('target_node_uuid')}"
|
||||
)
|
||||
if chain not in seen_chains:
|
||||
seen_chains.add(chain)
|
||||
relationship_chains.append(chain)
|
||||
return LocalInsightForgeResult(
|
||||
query=query,
|
||||
simulation_requirement=simulation_requirement,
|
||||
sub_queries=[query] if query else [],
|
||||
semantic_facts=facts,
|
||||
entity_insights=entity_insights,
|
||||
relationship_chains=relationship_chains,
|
||||
total_facts=len(facts),
|
||||
total_entities=len(entity_insights),
|
||||
total_relationships=len(relationship_chains),
|
||||
)
|
||||
|
||||
def panorama_search(
|
||||
self,
|
||||
*,
|
||||
graph_id: str,
|
||||
query: str,
|
||||
include_expired: bool = True,
|
||||
limit: int = 50,
|
||||
) -> LocalPanoramaResult:
|
||||
self._validate_graph_id(graph_id)
|
||||
safe_limit = min(max(int(limit), 1), 100)
|
||||
nodes = self.get_all_nodes(limit=10_000)
|
||||
edges = self.get_all_edges(limit=20_000)
|
||||
|
||||
active_facts = []
|
||||
historical_facts = []
|
||||
for edge in edges:
|
||||
if not edge.fact:
|
||||
continue
|
||||
if edge.is_invalid or edge.is_expired:
|
||||
valid_at = edge.valid_at or "Unknown"
|
||||
invalid_at = edge.invalid_at or edge.expired_at or "Unknown"
|
||||
historical_facts.append(f"[{valid_at} - {invalid_at}] {edge.fact}")
|
||||
else:
|
||||
active_facts.append(edge.fact)
|
||||
|
||||
query_lower = query.casefold()
|
||||
keywords = [
|
||||
word.strip()
|
||||
for word in query_lower.replace(",", " ").replace(",", " ").split()
|
||||
if len(word.strip()) > 1
|
||||
]
|
||||
|
||||
def relevance_score(fact: str) -> int:
|
||||
fact_lower = fact.casefold()
|
||||
score = 100 if query_lower in fact_lower else 0
|
||||
return score + sum(10 for keyword in keywords if keyword in fact_lower)
|
||||
|
||||
active_facts.sort(key=relevance_score, reverse=True)
|
||||
historical_facts.sort(key=relevance_score, reverse=True)
|
||||
return LocalPanoramaResult(
|
||||
query=query,
|
||||
all_nodes=nodes,
|
||||
all_edges=edges,
|
||||
active_facts=active_facts[:safe_limit],
|
||||
historical_facts=historical_facts[:safe_limit] if include_expired else [],
|
||||
)
|
||||
|
||||
def get_entity_summary(self, *, graph_id: str, entity_name: str) -> dict[str, Any]:
|
||||
self._validate_graph_id(graph_id)
|
||||
normalized = " ".join((entity_name or "").casefold().split())
|
||||
node = next(
|
||||
(
|
||||
item
|
||||
for item in self.repository.list_nodes(limit=10_000)
|
||||
if item.normalized_name == normalized
|
||||
),
|
||||
None,
|
||||
)
|
||||
if node is None:
|
||||
return {"entity_name": entity_name, "found": False, "summary": "", "related_facts": []}
|
||||
edges = self.get_node_edges(node.id, limit=500)
|
||||
return {
|
||||
"entity_name": node.canonical_name,
|
||||
"found": True,
|
||||
"uuid": node.id,
|
||||
"labels": list(node.labels or []),
|
||||
"summary": node.summary or "",
|
||||
"attributes": dict(node.attributes or {}),
|
||||
"related_facts": [edge.fact for edge in edges],
|
||||
}
|
||||
|
||||
def get_entities_by_type(self, *, graph_id: str, entity_type: str) -> list[LocalNodeInfo]:
|
||||
self._validate_graph_id(graph_id)
|
||||
target = (entity_type or "").casefold().strip()
|
||||
return [
|
||||
self._node_info(node)
|
||||
for node in self.repository.list_nodes(limit=10_000)
|
||||
if target in {str(label).casefold() for label in (node.labels or [])}
|
||||
]
|
||||
|
||||
def interview_agents(
|
||||
self,
|
||||
*,
|
||||
simulation_id: str,
|
||||
interview_requirement: str,
|
||||
simulation_requirement: str,
|
||||
max_agents: int = 5,
|
||||
) -> LocalInterviewResult:
|
||||
del simulation_id, max_agents
|
||||
return LocalInterviewResult(
|
||||
interview_topic=interview_requirement,
|
||||
summary=(
|
||||
"Local memory stores graph facts, not simulation transcripts. "
|
||||
"Use the retrieved entities and relationships as evidence; no synthetic interview response was generated."
|
||||
),
|
||||
)
|
||||
|
||||
def get_simulation_context(
|
||||
self,
|
||||
*,
|
||||
graph_id: str,
|
||||
simulation_requirement: str,
|
||||
limit: int = 30,
|
||||
) -> dict[str, Any]:
|
||||
self._validate_graph_id(graph_id)
|
||||
result = self.search_graph(
|
||||
simulation_requirement,
|
||||
limit=limit,
|
||||
scope="both",
|
||||
graph_id=graph_id,
|
||||
)
|
||||
nodes = self.repository.list_nodes(limit=10_000)
|
||||
entities = [
|
||||
{
|
||||
"name": node.canonical_name,
|
||||
"type": next((label for label in (node.labels or []) if label not in {"Entity", "Node"}), "Entity"),
|
||||
"summary": node.summary or "",
|
||||
}
|
||||
for node in nodes
|
||||
if any(label not in {"Entity", "Node"} for label in (node.labels or []))
|
||||
]
|
||||
return {
|
||||
"simulation_requirement": simulation_requirement,
|
||||
"related_facts": result.facts,
|
||||
"graph_statistics": self.get_graph_statistics(graph_id),
|
||||
"entities": entities[: max(int(limit), 1)],
|
||||
"total_entities": len(entities),
|
||||
}
|
||||
|
||||
def get_graph_statistics(self, graph_id: str | None = None) -> dict[str, Any]:
|
||||
self._validate_graph_id(graph_id)
|
||||
nodes = self.repository.list_nodes(limit=10_000)
|
||||
edges = self.repository.list_edges(limit=20_000)
|
||||
entity_types: dict[str, int] = {}
|
||||
for node in nodes:
|
||||
for label in node.labels or []:
|
||||
if label not in {"Entity", "Node"}:
|
||||
entity_types[label] = entity_types.get(label, 0) + 1
|
||||
relation_types: dict[str, int] = {}
|
||||
for edge in edges:
|
||||
relation_types[edge.relation] = relation_types.get(edge.relation, 0) + 1
|
||||
return {
|
||||
"graph_id": self.repository.graph_id,
|
||||
"node_count": len(nodes),
|
||||
"edge_count": len(edges),
|
||||
"total_nodes": len(nodes),
|
||||
"total_edges": len(edges),
|
||||
"entity_types": entity_types,
|
||||
"relation_types": relation_types,
|
||||
}
|
||||
@@ -8,6 +8,8 @@ OASIS Agent Profile生成器
|
||||
3. 区分个人实体和抽象群体实体
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import random
|
||||
import time
|
||||
@@ -16,12 +18,13 @@ from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
|
||||
from openai import OpenAI
|
||||
from zep_cloud.client import Zep
|
||||
from typing import TYPE_CHECKING
|
||||
|
||||
from ..config import Config
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import get_language_instruction, get_locale, set_locale, t
|
||||
from .zep_entity_reader import EntityNode, ZepEntityReader
|
||||
if TYPE_CHECKING:
|
||||
from .zep_entity_reader import EntityNode
|
||||
|
||||
logger = get_logger('crowdsight.oasis_profile')
|
||||
|
||||
@@ -184,7 +187,9 @@ class OasisProfileGenerator:
|
||||
base_url: Optional[str] = None,
|
||||
model_name: Optional[str] = None,
|
||||
zep_api_key: Optional[str] = None,
|
||||
graph_id: Optional[str] = None
|
||||
graph_id: Optional[str] = None,
|
||||
use_zep_context: bool = True,
|
||||
local_memory_tools: Optional[Any] = None,
|
||||
):
|
||||
self.api_key = api_key or Config.LLM_API_KEY
|
||||
self.base_url = base_url or Config.LLM_BASE_URL
|
||||
@@ -198,13 +203,18 @@ class OasisProfileGenerator:
|
||||
base_url=self.base_url
|
||||
)
|
||||
|
||||
# Zep客户端用于检索丰富上下文
|
||||
self.zep_api_key = zep_api_key or Config.ZEP_API_KEY
|
||||
# Local mode is a hard boundary: never construct a Zep client, even
|
||||
# when a legacy Zep key is present in the process environment.
|
||||
zep_context_enabled = use_zep_context and Config.MEMORY_BACKEND != "local"
|
||||
self.zep_api_key = (zep_api_key or Config.ZEP_API_KEY) if zep_context_enabled else None
|
||||
self.zep_client = None
|
||||
self.graph_id = graph_id
|
||||
self.local_memory_tools = local_memory_tools
|
||||
|
||||
if self.zep_api_key:
|
||||
try:
|
||||
from zep_cloud.client import Zep
|
||||
|
||||
self.zep_client = Zep(api_key=self.zep_api_key)
|
||||
except Exception as e:
|
||||
logger.warning(f"Zep客户端初始化失败: {e}")
|
||||
@@ -285,6 +295,45 @@ class OasisProfileGenerator:
|
||||
suffix = random.randint(100, 999)
|
||||
return f"{username}_{suffix}"
|
||||
|
||||
def _search_local_for_entity(self, entity: EntityNode) -> Dict[str, Any]:
|
||||
"""Retrieve bounded deterministic context from the local memory tools."""
|
||||
empty = {"facts": [], "node_summaries": [], "context": ""}
|
||||
if self.local_memory_tools is None:
|
||||
return empty
|
||||
try:
|
||||
result = self.local_memory_tools.search_graph(
|
||||
entity.name,
|
||||
limit=30,
|
||||
scope="both",
|
||||
)
|
||||
facts = list(getattr(result, "facts", []) or [])
|
||||
node_summaries = []
|
||||
for node in list(getattr(result, "nodes", []) or []):
|
||||
if not isinstance(node, dict):
|
||||
continue
|
||||
summary = node.get("summary")
|
||||
name = node.get("name")
|
||||
if summary:
|
||||
node_summaries.append(str(summary))
|
||||
elif name and name != entity.name:
|
||||
node_summaries.append(f"Related entity: {name}")
|
||||
context_parts = []
|
||||
if facts:
|
||||
context_parts.append("### Local memory facts\n" + "\n".join(f"- {fact}" for fact in facts[:20]))
|
||||
if node_summaries:
|
||||
context_parts.append(
|
||||
"### Local memory related entities\n"
|
||||
+ "\n".join(f"- {summary}" for summary in node_summaries[:10])
|
||||
)
|
||||
return {
|
||||
"facts": facts,
|
||||
"node_summaries": node_summaries,
|
||||
"context": "\n\n".join(context_parts),
|
||||
}
|
||||
except Exception as exc:
|
||||
logger.warning("Local memory context lookup failed: %s", type(exc).__name__)
|
||||
return empty
|
||||
|
||||
def _search_zep_for_entity(self, entity: EntityNode) -> Dict[str, Any]:
|
||||
"""
|
||||
使用Zep图谱混合搜索功能获取实体相关的丰富信息
|
||||
@@ -474,17 +523,28 @@ class OasisProfileGenerator:
|
||||
if related_info:
|
||||
context_parts.append("### 关联实体信息\n" + "\n".join(related_info))
|
||||
|
||||
# 4. 使用Zep混合检索获取更丰富的信息
|
||||
zep_results = self._search_zep_for_entity(entity)
|
||||
|
||||
if zep_results.get("facts"):
|
||||
# Use exactly one deterministic memory backend for enrichment.
|
||||
if self.local_memory_tools is not None:
|
||||
memory_results = self._search_local_for_entity(entity)
|
||||
memory_label = "Local memory"
|
||||
else:
|
||||
memory_results = self._search_zep_for_entity(entity)
|
||||
memory_label = "Zep"
|
||||
|
||||
if memory_results.get("facts"):
|
||||
# 去重:排除已存在的事实
|
||||
new_facts = [f for f in zep_results["facts"] if f not in existing_facts]
|
||||
new_facts = [f for f in memory_results["facts"] if f not in existing_facts]
|
||||
if new_facts:
|
||||
context_parts.append("### Zep检索到的事实信息\n" + "\n".join(f"- {f}" for f in new_facts[:15]))
|
||||
|
||||
if zep_results.get("node_summaries"):
|
||||
context_parts.append("### Zep检索到的相关节点\n" + "\n".join(f"- {s}" for s in zep_results["node_summaries"][:10]))
|
||||
context_parts.append(
|
||||
f"### {memory_label} retrieved facts\n"
|
||||
+ "\n".join(f"- {f}" for f in new_facts[:15])
|
||||
)
|
||||
|
||||
if memory_results.get("node_summaries"):
|
||||
context_parts.append(
|
||||
f"### {memory_label} related nodes\n"
|
||||
+ "\n".join(f"- {s}" for s in memory_results["node_summaries"][:10])
|
||||
)
|
||||
|
||||
return "\n\n".join(context_parts)
|
||||
|
||||
|
||||
76
backend/app/services/password_reset.py
Normal file
76
backend/app/services/password_reset.py
Normal file
@@ -0,0 +1,76 @@
|
||||
"""Durable, single-use password reset token service.
|
||||
|
||||
Tokens are opaque random strings; only their SHA-256 hash is stored. Consuming
|
||||
a valid token marks it used and bumps the user's ``auth_version`` so previously
|
||||
issued sessions are invalidated.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import hashlib
|
||||
import secrets
|
||||
from datetime import datetime, timedelta, timezone
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.password_reset import PasswordResetToken
|
||||
from ..models.saas import User
|
||||
|
||||
DEFAULT_TTL = timedelta(hours=24)
|
||||
DEFAULT_TTL_HOURS = 24
|
||||
|
||||
|
||||
def _hash(token: str) -> str:
|
||||
return hashlib.sha256(token.encode("utf-8")).hexdigest()
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
# Use naive UTC: SQLite returns naive datetimes for DateTime(timezone=True)
|
||||
# on read, so comparisons are kept tz-naive to avoid aware/naive errors.
|
||||
return datetime.utcnow()
|
||||
|
||||
|
||||
class PasswordResetService:
|
||||
def __init__(self, session: Session, *, ttl: timedelta | None = None):
|
||||
self.session = session
|
||||
self.ttl = ttl or DEFAULT_TTL
|
||||
|
||||
def create_token(self, *, user_id: str) -> str:
|
||||
token = secrets.token_urlsafe(32)
|
||||
record = PasswordResetToken(
|
||||
user_id=user_id,
|
||||
token_hash=_hash(token),
|
||||
expires_at=_utc_now() + self.ttl,
|
||||
)
|
||||
self.session.add(record)
|
||||
self.session.flush()
|
||||
return token
|
||||
|
||||
def _latest_token(self, user_id: str) -> Optional[PasswordResetToken]:
|
||||
return (
|
||||
self.session.query(PasswordResetToken)
|
||||
.filter(PasswordResetToken.user_id == user_id)
|
||||
.order_by(PasswordResetToken.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
def consume_token(self, token: str, *, user_id: str) -> bool:
|
||||
"""Consume a valid, unexpired, unused token for ``user_id``.
|
||||
|
||||
On success marks the token used and increments the user's auth_version
|
||||
(invalidating older sessions). Returns False otherwise.
|
||||
"""
|
||||
record = self._latest_token(user_id)
|
||||
if record is None or record.used:
|
||||
return False
|
||||
if _hash(token) != record.token_hash:
|
||||
return False
|
||||
if record.expires_at < _utc_now():
|
||||
return False
|
||||
record.used = True
|
||||
user = self.session.get(User, user_id)
|
||||
if user is not None:
|
||||
user.auth_version = (user.auth_version or 0) + 1
|
||||
self.session.flush()
|
||||
return True
|
||||
379
backend/app/services/product_repository.py
Normal file
379
backend/app/services/product_repository.py
Normal file
@@ -0,0 +1,379 @@
|
||||
"""Tenant-scoped repository for durable product resources.
|
||||
|
||||
Projects, simulations, and reports owned by an organization and (for projects)
|
||||
an owner/creator. Every lookup requires an explicit ``organization_id`` so an
|
||||
unscoped id supplied by a route can never be read across tenant boundaries.
|
||||
|
||||
This repository is flush-only: the caller owns transaction boundaries and must
|
||||
commit. Nothing here persists secrets or raw exception details.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
|
||||
class ProductRepository:
|
||||
"""Flush-only durable read/write for product resources."""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
# ---- projects ----
|
||||
|
||||
def create_project(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
owner_user_id: str | None,
|
||||
name: str,
|
||||
language: str = "en",
|
||||
ontology: dict | None = None,
|
||||
):
|
||||
from app.models.product import ProductProject
|
||||
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
raise ValueError("organization_id_required")
|
||||
if not isinstance(name, str):
|
||||
raise ValueError("invalid_project_name")
|
||||
project = ProductProject(
|
||||
organization_id=organization_id,
|
||||
owner_user_id=owner_user_id,
|
||||
name=name,
|
||||
language=language,
|
||||
ontology=ontology,
|
||||
)
|
||||
self.session.add(project)
|
||||
self.session.flush()
|
||||
return project
|
||||
|
||||
def get_project(self, project_id: str, *, organization_id: str):
|
||||
from app.models.product import ProductProject
|
||||
|
||||
return (
|
||||
self.session.query(ProductProject)
|
||||
.filter(
|
||||
ProductProject.id == project_id,
|
||||
ProductProject.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
def sync_project(self, state, *, commit: bool = False):
|
||||
"""Mirror a legacy filesystem Project state into durable SQL (idempotent).
|
||||
|
||||
``state`` may be a ``Project`` object or any object exposing the durable
|
||||
fields (e.g. ``ProjectManager.to_dict()``). The sync is an upsert keyed on
|
||||
the project id + organization so repeated writes do not duplicate rows.
|
||||
"""
|
||||
from app.models.product import ProductProject, ProjectStatus
|
||||
|
||||
def _get(names, default=None):
|
||||
for name in names:
|
||||
value = getattr(state, name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
if isinstance(state, dict):
|
||||
value = state.get(name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return default
|
||||
|
||||
project_id = _get(["project_id", "id"])
|
||||
organization_id = _get(["organization_id"])
|
||||
if not project_id or not organization_id:
|
||||
raise ValueError("project_sync_requires_id_and_organization")
|
||||
|
||||
row = (
|
||||
self.session.query(ProductProject)
|
||||
.filter(
|
||||
ProductProject.id == project_id,
|
||||
ProductProject.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
row = ProductProject(
|
||||
id=project_id,
|
||||
organization_id=organization_id,
|
||||
owner_user_id=_get(["owner_user_id"]),
|
||||
name=_get(["name"], "") or "",
|
||||
status=_get(["status"], ProjectStatus.CREATED.value),
|
||||
language=_get(["language"], "en"),
|
||||
total_text_length=int(_get(["total_text_length"], 0) or 0),
|
||||
source_metadata=_get(["source_metadata"]),
|
||||
ontology=_get(["ontology"]),
|
||||
analysis_summary=_get(["analysis_summary"]),
|
||||
simulation_requirement=_get(["simulation_requirement"]),
|
||||
graph_id=_get(["graph_id"]),
|
||||
graph_build_task_id=_get(["graph_build_task_id"]),
|
||||
error=_get(["error"]),
|
||||
)
|
||||
self.session.add(row)
|
||||
else:
|
||||
row.owner_user_id = _get(["owner_user_id"], row.owner_user_id)
|
||||
row.name = _get(["name"], row.name) or row.name
|
||||
row.status = _get(["status"], row.status)
|
||||
row.language = _get(["language"], row.language)
|
||||
row.total_text_length = int(_get(["total_text_length"], row.total_text_length) or 0)
|
||||
if (_get(["ontology"])) is not None:
|
||||
row.ontology = _get(["ontology"])
|
||||
if (_get(["analysis_summary"])) is not None:
|
||||
row.analysis_summary = _get(["analysis_summary"])
|
||||
if (_get(["simulation_requirement"])) is not None:
|
||||
row.simulation_requirement = _get(["simulation_requirement"])
|
||||
if (_get(["graph_id"])) is not None:
|
||||
row.graph_id = _get(["graph_id"])
|
||||
if (_get(["graph_build_task_id"])) is not None:
|
||||
row.graph_build_task_id = _get(["graph_build_task_id"])
|
||||
if (_get(["error"])) is not None:
|
||||
row.error = _get(["error"])
|
||||
self.session.flush()
|
||||
if commit:
|
||||
self.session.commit()
|
||||
return row
|
||||
|
||||
def list_projects(
|
||||
self, *, organization_id: str, owner_user_id: str | None = None, limit: int = 200
|
||||
) -> list:
|
||||
from app.models.product import ProductProject
|
||||
|
||||
query = self.session.query(ProductProject).filter(
|
||||
ProductProject.organization_id == organization_id
|
||||
)
|
||||
if owner_user_id is not None:
|
||||
query = query.filter(ProductProject.owner_user_id == owner_user_id)
|
||||
return query.order_by(ProductProject.created_at.desc()).limit(limit).all()
|
||||
|
||||
# ---- simulations ----
|
||||
|
||||
def create_simulation(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
project_id: str,
|
||||
created_by_user_id: str | None,
|
||||
status: str = "created",
|
||||
platform: str = "parallel",
|
||||
config: dict | None = None,
|
||||
):
|
||||
from app.models.product import ProductSimulation
|
||||
|
||||
self._require_org(organization_id)
|
||||
simulation = ProductSimulation(
|
||||
organization_id=organization_id,
|
||||
project_id=project_id,
|
||||
created_by_user_id=created_by_user_id,
|
||||
status=status,
|
||||
platform=platform,
|
||||
config=config,
|
||||
)
|
||||
self.session.add(simulation)
|
||||
self.session.flush()
|
||||
return simulation
|
||||
|
||||
def sync_simulation(self, state, *, commit: bool = False):
|
||||
"""Mirror a legacy simulation state into durable SQL (idempotent upsert)."""
|
||||
from app.models.product import ProductSimulation
|
||||
|
||||
def _get(names, default=None):
|
||||
for name in names:
|
||||
value = getattr(state, name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
if isinstance(state, dict):
|
||||
value = state.get(name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return default
|
||||
|
||||
simulation_id = _get(["simulation_id", "id"])
|
||||
organization_id = _get(["organization_id"])
|
||||
project_id = _get(["project_id"])
|
||||
if not simulation_id or not organization_id:
|
||||
raise ValueError("simulation_sync_requires_id_and_organization")
|
||||
|
||||
row = (
|
||||
self.session.query(ProductSimulation)
|
||||
.filter(
|
||||
ProductSimulation.id == simulation_id,
|
||||
ProductSimulation.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
row = ProductSimulation(
|
||||
id=simulation_id,
|
||||
organization_id=organization_id,
|
||||
project_id=project_id or "",
|
||||
created_by_user_id=_get(["created_by_user_id", "owner_user_id"]),
|
||||
status=_get(["status"], "created"),
|
||||
platform=_get(["platform"], "parallel"),
|
||||
config=_get(["config"]),
|
||||
current_round=int(_get(["current_round"], 0) or 0),
|
||||
)
|
||||
self.session.add(row)
|
||||
else:
|
||||
row.project_id = project_id or row.project_id
|
||||
row.status = _get(["status"], row.status)
|
||||
row.platform = _get(["platform"], row.platform)
|
||||
if (_get(["config"])) is not None:
|
||||
row.config = _get(["config"])
|
||||
row.current_round = int(_get(["current_round"], row.current_round) or 0)
|
||||
self.session.flush()
|
||||
if commit:
|
||||
self.session.commit()
|
||||
return row
|
||||
|
||||
def get_simulation(self, simulation_id: str, *, organization_id: str):
|
||||
from app.models.product import ProductSimulation
|
||||
|
||||
return (
|
||||
self.session.query(ProductSimulation)
|
||||
.filter(
|
||||
ProductSimulation.id == simulation_id,
|
||||
ProductSimulation.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
def list_simulations(
|
||||
self, *, organization_id: str, project_id: str | None = None, limit: int = 200
|
||||
) -> list:
|
||||
from app.models.product import ProductSimulation
|
||||
|
||||
query = self.session.query(ProductSimulation).filter(
|
||||
ProductSimulation.organization_id == organization_id
|
||||
)
|
||||
if project_id is not None:
|
||||
query = query.filter(ProductSimulation.project_id == project_id)
|
||||
return query.order_by(ProductSimulation.created_at.desc()).limit(limit).all()
|
||||
|
||||
# ---- reports ----
|
||||
|
||||
def create_report(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
project_id: str,
|
||||
simulation_id: str | None,
|
||||
created_by_user_id: str | None,
|
||||
title: str = "",
|
||||
status: str = "draft",
|
||||
):
|
||||
from app.models.product import DurableReport
|
||||
|
||||
self._require_org(organization_id)
|
||||
report = DurableReport(
|
||||
organization_id=organization_id,
|
||||
project_id=project_id,
|
||||
simulation_id=simulation_id,
|
||||
created_by_user_id=created_by_user_id,
|
||||
title=title,
|
||||
status=status,
|
||||
)
|
||||
self.session.add(report)
|
||||
self.session.flush()
|
||||
return report
|
||||
|
||||
def get_report(self, report_id: str, *, organization_id: str):
|
||||
from app.models.product import DurableReport
|
||||
|
||||
return (
|
||||
self.session.query(DurableReport)
|
||||
.filter(
|
||||
DurableReport.id == report_id,
|
||||
DurableReport.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
|
||||
def sync_report(
|
||||
self,
|
||||
state,
|
||||
*,
|
||||
organization_id: str,
|
||||
project_id: str,
|
||||
simulation_id: str | None,
|
||||
created_by_user_id: str | None,
|
||||
commit: bool = False,
|
||||
):
|
||||
"""Mirror a legacy report state into durable SQL (idempotent upsert)."""
|
||||
from app.models.product import DurableReport
|
||||
|
||||
def _get(names, default=None):
|
||||
for name in names:
|
||||
value = getattr(state, name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
if isinstance(state, dict):
|
||||
value = state.get(name, None)
|
||||
if value is not None:
|
||||
return value
|
||||
return default
|
||||
|
||||
report_id = _get(["report_id", "id"])
|
||||
if not report_id:
|
||||
raise ValueError("report_sync_requires_id")
|
||||
|
||||
row = (
|
||||
self.session.query(DurableReport)
|
||||
.filter(
|
||||
DurableReport.id == report_id,
|
||||
DurableReport.organization_id == organization_id,
|
||||
)
|
||||
.first()
|
||||
)
|
||||
if row is None:
|
||||
row = DurableReport(
|
||||
id=report_id,
|
||||
organization_id=organization_id,
|
||||
project_id=project_id or "",
|
||||
simulation_id=simulation_id,
|
||||
created_by_user_id=created_by_user_id,
|
||||
status=_get(["status"], "draft"),
|
||||
title=_get(["title"], "") or "",
|
||||
outline=_get(["outline"]),
|
||||
markdown_content=_get(["markdown_content", "content"]),
|
||||
error=_get(["error"]),
|
||||
)
|
||||
self.session.add(row)
|
||||
else:
|
||||
row.status = _get(["status"], row.status)
|
||||
_title = _get(["title"])
|
||||
if _title is not None:
|
||||
row.title = _title
|
||||
if (_get(["outline"])) is not None:
|
||||
row.outline = _get(["outline"])
|
||||
if (_get(["markdown_content", "content"])) is not None:
|
||||
row.markdown_content = _get(["markdown_content", "content"])
|
||||
if (_get(["error"])) is not None:
|
||||
row.error = _get(["error"])
|
||||
self.session.flush()
|
||||
if commit:
|
||||
self.session.commit()
|
||||
return row
|
||||
|
||||
def list_reports(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
project_id: str | None = None,
|
||||
simulation_id: str | None = None,
|
||||
limit: int = 200,
|
||||
) -> list:
|
||||
from app.models.product import DurableReport
|
||||
|
||||
query = self.session.query(DurableReport).filter(
|
||||
DurableReport.organization_id == organization_id
|
||||
)
|
||||
if project_id is not None:
|
||||
query = query.filter(DurableReport.project_id == project_id)
|
||||
if simulation_id is not None:
|
||||
query = query.filter(DurableReport.simulation_id == simulation_id)
|
||||
return query.order_by(DurableReport.created_at.desc()).limit(limit).all()
|
||||
|
||||
@staticmethod
|
||||
def _require_org(organization_id: str) -> None:
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
raise ValueError("organization_id_required")
|
||||
49
backend/app/services/rate_limiter.py
Normal file
49
backend/app/services/rate_limiter.py
Normal file
@@ -0,0 +1,49 @@
|
||||
"""Durable rate limiting (sliding-window counter) over the rate_limit_events table.
|
||||
|
||||
A durable counter table means limits survive worker restarts and multi-instance
|
||||
deploys. Keys are typically ``<operation>:<email-or-ip>``. No secrets are stored.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.rate_limit import RateLimitEvent
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class RateLimiter:
|
||||
"""Flush-only rate limiter; caller owns the transaction boundary."""
|
||||
|
||||
def __init__(self, session: Session, *, window: timedelta | None = None, limit: int | None = None):
|
||||
self.session = session
|
||||
self.window = window or timedelta(minutes=15)
|
||||
self.limit = limit or 60
|
||||
|
||||
def check_and_record(self, operation: str, *, key: str, organization_id: str | None = None) -> bool:
|
||||
"""Record a hit if under the limit; return False when over the limit."""
|
||||
cutoff = _utc_now() - self.window
|
||||
count = self.session.query(RateLimitEvent).filter(
|
||||
RateLimitEvent.operation == operation,
|
||||
RateLimitEvent.key == key,
|
||||
RateLimitEvent.created_at >= cutoff,
|
||||
).count()
|
||||
if count >= self.limit:
|
||||
return False
|
||||
self.session.add(
|
||||
RateLimitEvent(
|
||||
operation=operation,
|
||||
key=key,
|
||||
organization_id=organization_id,
|
||||
)
|
||||
)
|
||||
# Rate-limit events are append-only side effects that must survive even
|
||||
# if the surrounding request transaction is later rolled back, so commit
|
||||
# them immediately.
|
||||
self.session.commit()
|
||||
return True
|
||||
@@ -22,13 +22,6 @@ from ..config import Config
|
||||
from ..utils.llm_client import LLMClient
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import get_language_instruction, t
|
||||
from .zep_tools import (
|
||||
ZepToolsService,
|
||||
SearchResult,
|
||||
InsightForgeResult,
|
||||
PanoramaResult,
|
||||
InterviewResult
|
||||
)
|
||||
|
||||
logger = get_logger('crowdsight.report_agent')
|
||||
|
||||
@@ -886,7 +879,8 @@ class ReportAgent:
|
||||
simulation_id: str,
|
||||
simulation_requirement: str,
|
||||
llm_client: Optional[LLMClient] = None,
|
||||
zep_tools: Optional[ZepToolsService] = None
|
||||
memory_tools: Optional[Any] = None,
|
||||
zep_tools: Optional[Any] = None,
|
||||
):
|
||||
"""
|
||||
初始化Report Agent
|
||||
@@ -903,7 +897,20 @@ class ReportAgent:
|
||||
self.simulation_requirement = simulation_requirement
|
||||
|
||||
self.llm = llm_client or LLMClient()
|
||||
self.zep_tools = zep_tools or ZepToolsService()
|
||||
if memory_tools is not None:
|
||||
self.memory_tools = memory_tools
|
||||
elif Config.MEMORY_BACKEND == "local":
|
||||
# Local mode must be explicit. Constructing a Zep client here would
|
||||
# create a silent backend fallback and can cross tenant boundaries.
|
||||
raise ValueError("memory_tools_required_for_local_backend")
|
||||
elif zep_tools is not None:
|
||||
# ``zep_tools`` remains a compatibility injection for the legacy
|
||||
# backend; it is never considered in local mode.
|
||||
self.memory_tools = zep_tools
|
||||
else:
|
||||
from .zep_tools import ZepToolsService
|
||||
|
||||
self.memory_tools = ZepToolsService()
|
||||
|
||||
# 工具定义
|
||||
self.tools = self._define_tools()
|
||||
@@ -970,7 +977,7 @@ class ReportAgent:
|
||||
if tool_name == "insight_forge":
|
||||
query = parameters.get("query", "")
|
||||
ctx = parameters.get("report_context", "") or report_context
|
||||
result = self.zep_tools.insight_forge(
|
||||
result = self.memory_tools.insight_forge(
|
||||
graph_id=self.graph_id,
|
||||
query=query,
|
||||
simulation_requirement=self.simulation_requirement,
|
||||
@@ -984,7 +991,7 @@ class ReportAgent:
|
||||
include_expired = parameters.get("include_expired", True)
|
||||
if isinstance(include_expired, str):
|
||||
include_expired = include_expired.lower() in ['true', '1', 'yes']
|
||||
result = self.zep_tools.panorama_search(
|
||||
result = self.memory_tools.panorama_search(
|
||||
graph_id=self.graph_id,
|
||||
query=query,
|
||||
include_expired=include_expired
|
||||
@@ -997,7 +1004,7 @@ class ReportAgent:
|
||||
limit = parameters.get("limit", 10)
|
||||
if isinstance(limit, str):
|
||||
limit = int(limit)
|
||||
result = self.zep_tools.quick_search(
|
||||
result = self.memory_tools.quick_search(
|
||||
graph_id=self.graph_id,
|
||||
query=query,
|
||||
limit=limit
|
||||
@@ -1011,7 +1018,7 @@ class ReportAgent:
|
||||
if isinstance(max_agents, str):
|
||||
max_agents = int(max_agents)
|
||||
max_agents = min(max_agents, 10)
|
||||
result = self.zep_tools.interview_agents(
|
||||
result = self.memory_tools.interview_agents(
|
||||
simulation_id=self.simulation_id,
|
||||
interview_requirement=interview_topic,
|
||||
simulation_requirement=self.simulation_requirement,
|
||||
@@ -1027,12 +1034,12 @@ class ReportAgent:
|
||||
return self._execute_tool("quick_search", parameters, report_context)
|
||||
|
||||
elif tool_name == "get_graph_statistics":
|
||||
result = self.zep_tools.get_graph_statistics(self.graph_id)
|
||||
result = self.memory_tools.get_graph_statistics(self.graph_id)
|
||||
return json.dumps(result, ensure_ascii=False, indent=2)
|
||||
|
||||
elif tool_name == "get_entity_summary":
|
||||
entity_name = parameters.get("entity_name", "")
|
||||
result = self.zep_tools.get_entity_summary(
|
||||
result = self.memory_tools.get_entity_summary(
|
||||
graph_id=self.graph_id,
|
||||
entity_name=entity_name
|
||||
)
|
||||
@@ -1046,7 +1053,7 @@ class ReportAgent:
|
||||
|
||||
elif tool_name == "get_entities_by_type":
|
||||
entity_type = parameters.get("entity_type", "")
|
||||
nodes = self.zep_tools.get_entities_by_type(
|
||||
nodes = self.memory_tools.get_entities_by_type(
|
||||
graph_id=self.graph_id,
|
||||
entity_type=entity_type
|
||||
)
|
||||
@@ -1056,9 +1063,9 @@ class ReportAgent:
|
||||
else:
|
||||
return f"Unknown tool: {tool_name}. Please use one of the following tools: insight_forge, panorama_search, quick_search"
|
||||
|
||||
except Exception as e:
|
||||
logger.error(t('report.toolExecFailed', toolName=tool_name, error=str(e)))
|
||||
return f"Tool execution failed: {str(e)}"
|
||||
except Exception as exc:
|
||||
logger.error("Report tool execution failed: tool=%s error=%s", tool_name, type(exc).__name__)
|
||||
return "Tool execution failed; no result is available."
|
||||
|
||||
# 合法的工具名称集合,用于裸 JSON 兜底解析时校验
|
||||
VALID_TOOL_NAMES = {"insight_forge", "panorama_search", "quick_search", "interview_agents"}
|
||||
@@ -1154,7 +1161,7 @@ class ReportAgent:
|
||||
progress_callback("planning", 0, t('progress.analyzingRequirements'))
|
||||
|
||||
# 首先获取模拟上下文
|
||||
context = self.zep_tools.get_simulation_context(
|
||||
context = self.memory_tools.get_simulation_context(
|
||||
graph_id=self.graph_id,
|
||||
simulation_requirement=self.simulation_requirement
|
||||
)
|
||||
@@ -1737,19 +1744,22 @@ class ReportAgent:
|
||||
return report
|
||||
|
||||
except Exception as e:
|
||||
logger.error(t('report.reportGenFailed', error=str(e)))
|
||||
logger.error(
|
||||
"Report generation failed: error_type=%s",
|
||||
type(e).__name__,
|
||||
)
|
||||
report.status = ReportStatus.FAILED
|
||||
report.error = str(e)
|
||||
report.error = t('api.internalError')
|
||||
|
||||
# 记录错误日志
|
||||
if self.report_logger:
|
||||
self.report_logger.log_error(str(e), "failed")
|
||||
self.report_logger.log_error("report_generation_failed", "failed")
|
||||
|
||||
# 保存失败状态
|
||||
try:
|
||||
ReportManager.save_report(report)
|
||||
ReportManager.update_progress(
|
||||
report_id, "failed", -1, t('progress.reportFailed', error=str(e)),
|
||||
report_id, "failed", -1, t('api.internalError'),
|
||||
completed_sections=completed_section_titles
|
||||
)
|
||||
except Exception:
|
||||
@@ -1797,7 +1807,10 @@ class ReportAgent:
|
||||
if len(report.markdown_content) > 15000:
|
||||
report_content += "\n\n... [报告内容已截断] ..."
|
||||
except Exception as e:
|
||||
logger.warning(t('report.fetchReportFailed', error=e))
|
||||
logger.warning(
|
||||
"Failed to fetch report for chat: error_type=%s",
|
||||
type(e).__name__,
|
||||
)
|
||||
|
||||
system_prompt = CHAT_SYSTEM_PROMPT_TEMPLATE.format(
|
||||
simulation_requirement=self.simulation_requirement,
|
||||
|
||||
115
backend/app/services/settings_service.py
Normal file
115
backend/app/services/settings_service.py
Normal file
@@ -0,0 +1,115 @@
|
||||
"""Versioned, redacted platform settings service.
|
||||
|
||||
The API key is encrypted with a Fernet key derived from ``SECRET_KEY`` so the
|
||||
plaintext never appears in the record, API responses, or logs. The public API
|
||||
surface only ever sees a masked value and the settings version; a per-job
|
||||
snapshot references the version rather than the secret.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import base64
|
||||
import hashlib
|
||||
from datetime import datetime, timezone
|
||||
from typing import Optional
|
||||
from uuid import uuid4
|
||||
|
||||
from cryptography.fernet import Fernet
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..config import Config
|
||||
from ..models.settings import PlatformSettings
|
||||
|
||||
MASK = "sk-••••••••"
|
||||
|
||||
|
||||
def _utc_now() -> datetime:
|
||||
return datetime.now(timezone.utc)
|
||||
|
||||
|
||||
class SettingsService:
|
||||
"""Flush-only settings repository; caller owns transactions."""
|
||||
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def _fernet(self) -> Fernet:
|
||||
secret = Config.SECRET_KEY
|
||||
if not secret:
|
||||
raise ValueError("settings_secret_key_required")
|
||||
digest = hashlib.sha256(secret.encode("utf-8")).digest()
|
||||
key = base64.urlsafe_b64encode(digest)
|
||||
return Fernet(key)
|
||||
|
||||
def _encrypt_secret(self, api_key: str) -> str:
|
||||
return self._fernet().encrypt(api_key.encode("utf-8")).decode("utf-8")
|
||||
|
||||
def _decrypt_secret(self, secret_ref: str) -> Optional[str]:
|
||||
if not secret_ref:
|
||||
return None
|
||||
try:
|
||||
return self._fernet().decrypt(secret_ref.encode("utf-8")).decode("utf-8")
|
||||
except Exception:
|
||||
return None
|
||||
|
||||
def _latest_row(self) -> Optional[PlatformSettings]:
|
||||
return (
|
||||
self.session.query(PlatformSettings)
|
||||
.order_by(PlatformSettings.created_at.desc())
|
||||
.first()
|
||||
)
|
||||
|
||||
def save_settings(
|
||||
self,
|
||||
settings: dict,
|
||||
*,
|
||||
api_key: Optional[str] = None,
|
||||
updated_by: Optional[str] = None,
|
||||
) -> str:
|
||||
"""Persist a new version, encrypting the API key if provided."""
|
||||
version = f"v{uuid4().hex[:12]}"
|
||||
# Clear active on all existing rows, then insert the new active version.
|
||||
for row in self.session.query(PlatformSettings).filter(
|
||||
PlatformSettings.active.is_(True)
|
||||
):
|
||||
row.active = False
|
||||
record = PlatformSettings(
|
||||
version=version,
|
||||
settings=settings,
|
||||
secret_ref=self._encrypt_secret(api_key) if api_key else None,
|
||||
updated_by_user_id=updated_by,
|
||||
active=True,
|
||||
)
|
||||
self.session.add(record)
|
||||
self.session.flush()
|
||||
return version
|
||||
|
||||
def _mask(self, settings: dict) -> dict:
|
||||
out = dict(settings or {})
|
||||
# A secret value should never be in the public dict; be defensive.
|
||||
for key in list(out.keys()):
|
||||
if "key" in key.lower() or "secret" in key.lower() or "token" in key.lower():
|
||||
out[key] = MASK
|
||||
return out
|
||||
|
||||
def active_settings(self) -> dict:
|
||||
row = self._latest_row()
|
||||
if row is None:
|
||||
return {"settings": {}, "version": None, "updated_by": None, "api_key": MASK}
|
||||
masked = self._mask(row.settings if isinstance(row.settings, dict) else {})
|
||||
return {
|
||||
"settings": masked,
|
||||
"version": row.version,
|
||||
"updated_by": row.updated_by_user_id,
|
||||
"api_key": MASK if row.secret_ref else MASK,
|
||||
}
|
||||
|
||||
def snapshot_for_job(self) -> dict:
|
||||
row = self._latest_row()
|
||||
if row is None:
|
||||
return {"settings": {}, "version": None, "settings_version": None}
|
||||
return {
|
||||
"settings": self._mask(row.settings if isinstance(row.settings, dict) else {}),
|
||||
"version": row.version,
|
||||
"settings_version": row.version,
|
||||
}
|
||||
@@ -10,9 +10,11 @@
|
||||
4. 生成平台配置
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import json
|
||||
import math
|
||||
from typing import Dict, Any, List, Optional, Callable
|
||||
from typing import Dict, Any, List, Optional, Callable, TYPE_CHECKING
|
||||
from dataclasses import dataclass, field, asdict
|
||||
from datetime import datetime
|
||||
|
||||
@@ -21,7 +23,8 @@ from openai import OpenAI
|
||||
from ..config import Config
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import get_language_instruction, t
|
||||
from .zep_entity_reader import EntityNode, ZepEntityReader
|
||||
if TYPE_CHECKING:
|
||||
from .zep_entity_reader import EntityNode
|
||||
|
||||
logger = get_logger('crowdsight.simulation_config')
|
||||
|
||||
|
||||
@@ -7,15 +7,15 @@ OASIS模拟管理器
|
||||
import os
|
||||
import json
|
||||
import shutil
|
||||
from typing import Dict, Any, List, Optional
|
||||
from typing import Callable, Dict, Any, List, Optional
|
||||
from dataclasses import dataclass, field
|
||||
from datetime import datetime
|
||||
from enum import Enum
|
||||
|
||||
from ..config import Config
|
||||
from ..utils.logger import get_logger
|
||||
from .zep_entity_reader import ZepEntityReader, FilteredEntities
|
||||
from .oasis_profile_generator import OasisProfileGenerator, OasisAgentProfile
|
||||
from .memory_tools import LocalMemoryTools
|
||||
from .simulation_config_generator import SimulationConfigGenerator, SimulationParameters
|
||||
from ..utils.locale import t
|
||||
|
||||
@@ -47,6 +47,10 @@ class SimulationState:
|
||||
project_id: str
|
||||
graph_id: str
|
||||
|
||||
# Tenant/owner scope (populated by the creating route for durable sync).
|
||||
organization_id: Optional[str] = None
|
||||
owner_user_id: Optional[str] = None
|
||||
|
||||
# 平台启用状态
|
||||
enable_twitter: bool = True
|
||||
enable_reddit: bool = True
|
||||
@@ -81,6 +85,8 @@ class SimulationState:
|
||||
"simulation_id": self.simulation_id,
|
||||
"project_id": self.project_id,
|
||||
"graph_id": self.graph_id,
|
||||
"organization_id": self.organization_id,
|
||||
"owner_user_id": self.owner_user_id,
|
||||
"enable_twitter": self.enable_twitter,
|
||||
"enable_reddit": self.enable_reddit,
|
||||
"status": self.status.value,
|
||||
@@ -129,12 +135,25 @@ class SimulationManager:
|
||||
'../../uploads/simulations'
|
||||
)
|
||||
|
||||
def __init__(self):
|
||||
# 确保目录存在
|
||||
def __init__(
|
||||
self,
|
||||
entity_reader_factory: Optional[Callable[[str], Any]] = None,
|
||||
session_factory: Optional[Callable[..., Any]] = None,
|
||||
):
|
||||
# Ensure the legacy filesystem cache exists while the durable job migration is in progress.
|
||||
os.makedirs(self.SIMULATION_DATA_DIR, exist_ok=True)
|
||||
|
||||
# 内存中的模拟状态缓存
|
||||
self._simulations: Dict[str, SimulationState] = {}
|
||||
self._entity_reader_factory = entity_reader_factory
|
||||
self._session_factory = session_factory
|
||||
|
||||
def create_entity_reader(self, graph_id: str):
|
||||
if self._entity_reader_factory is not None:
|
||||
return self._entity_reader_factory(graph_id)
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
raise ValueError("local_entity_reader_factory_required")
|
||||
from .zep_entity_reader import ZepEntityReader
|
||||
|
||||
return ZepEntityReader()
|
||||
|
||||
def _get_simulation_dir(self, simulation_id: str) -> str:
|
||||
"""获取模拟数据目录"""
|
||||
@@ -153,6 +172,29 @@ class SimulationManager:
|
||||
json.dump(state.to_dict(), f, ensure_ascii=False, indent=2)
|
||||
|
||||
self._simulations[state.simulation_id] = state
|
||||
self._sync_simulation_to_durable(state)
|
||||
|
||||
def _sync_simulation_to_durable(self, state: SimulationState):
|
||||
"""Best-effort mirror of a simulation state into the durable table.
|
||||
|
||||
Only runs when tenant scope is known and a session factory is available;
|
||||
never raises so the filesystem manager remains authoritative during the
|
||||
migration.
|
||||
"""
|
||||
if not state.organization_id or not callable(self._session_factory):
|
||||
return
|
||||
try:
|
||||
session = self._session_factory()
|
||||
try:
|
||||
from .product_repository import ProductRepository
|
||||
|
||||
payload = state.to_dict()
|
||||
payload["created_by_user_id"] = state.owner_user_id
|
||||
ProductRepository(session).sync_simulation(payload, commit=True)
|
||||
finally:
|
||||
session.close()
|
||||
except Exception:
|
||||
logger.debug("durable simulation sync skipped", exc_info=True)
|
||||
|
||||
def _load_simulation_state(self, simulation_id: str) -> Optional[SimulationState]:
|
||||
"""从文件加载模拟状态"""
|
||||
@@ -172,6 +214,8 @@ class SimulationManager:
|
||||
simulation_id=simulation_id,
|
||||
project_id=data.get("project_id", ""),
|
||||
graph_id=data.get("graph_id", ""),
|
||||
organization_id=data.get("organization_id"),
|
||||
owner_user_id=data.get("owner_user_id"),
|
||||
enable_twitter=data.get("enable_twitter", True),
|
||||
enable_reddit=data.get("enable_reddit", True),
|
||||
status=SimulationStatus(data.get("status", "created")),
|
||||
@@ -271,18 +315,24 @@ class SimulationManager:
|
||||
|
||||
# ========== 阶段1: 读取并过滤实体 ==========
|
||||
if progress_callback:
|
||||
progress_callback("reading", 0, t('progress.connectingZepGraph'))
|
||||
|
||||
reader = ZepEntityReader()
|
||||
progress_callback("reading", 0, t('progress.connectingGraph' if Config.MEMORY_BACKEND == "local" else 'progress.connectingZepGraph'))
|
||||
|
||||
reader = self.create_entity_reader(state.graph_id)
|
||||
|
||||
if progress_callback:
|
||||
progress_callback("reading", 30, t('progress.readingNodeData'))
|
||||
|
||||
filtered = reader.filter_defined_entities(
|
||||
graph_id=state.graph_id,
|
||||
defined_entity_types=defined_entity_types,
|
||||
enrich_with_edges=True
|
||||
)
|
||||
|
||||
try:
|
||||
filtered = reader.filter_defined_entities(
|
||||
graph_id=state.graph_id,
|
||||
defined_entity_types=defined_entity_types,
|
||||
enrich_with_edges=True,
|
||||
)
|
||||
except Exception:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
raise
|
||||
|
||||
state.entities_count = filtered.filtered_count
|
||||
state.entity_types = list(filtered.entity_types)
|
||||
@@ -296,6 +346,9 @@ class SimulationManager:
|
||||
)
|
||||
|
||||
if filtered.filtered_count == 0:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
state.status = SimulationStatus.FAILED
|
||||
state.error = "没有找到符合条件的实体,请检查图谱是否正确构建"
|
||||
self._save_simulation_state(state)
|
||||
@@ -312,8 +365,18 @@ class SimulationManager:
|
||||
total=total_entities
|
||||
)
|
||||
|
||||
# 传入graph_id以启用Zep检索功能,获取更丰富的上下文
|
||||
generator = OasisProfileGenerator(graph_id=state.graph_id)
|
||||
local_memory_tools = None
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
repository = getattr(reader, "repository", None)
|
||||
if repository is None:
|
||||
raise ValueError("local_memory_reader_required")
|
||||
local_memory_tools = LocalMemoryTools(repository)
|
||||
|
||||
generator = OasisProfileGenerator(
|
||||
graph_id=state.graph_id,
|
||||
use_zep_context=Config.MEMORY_BACKEND != "local",
|
||||
local_memory_tools=local_memory_tools,
|
||||
)
|
||||
|
||||
def profile_progress(current, total, msg):
|
||||
if progress_callback:
|
||||
@@ -336,15 +399,20 @@ class SimulationManager:
|
||||
realtime_output_path = os.path.join(sim_dir, "twitter_profiles.csv")
|
||||
realtime_platform = "twitter"
|
||||
|
||||
profiles = generator.generate_profiles_from_entities(
|
||||
entities=filtered.entities,
|
||||
use_llm=use_llm_for_profiles,
|
||||
progress_callback=profile_progress,
|
||||
graph_id=state.graph_id, # 传入graph_id用于Zep检索
|
||||
parallel_count=parallel_profile_count, # 并行生成数量
|
||||
realtime_output_path=realtime_output_path, # 实时保存路径
|
||||
output_platform=realtime_platform # 输出格式
|
||||
)
|
||||
try:
|
||||
profiles = generator.generate_profiles_from_entities(
|
||||
entities=filtered.entities,
|
||||
use_llm=use_llm_for_profiles,
|
||||
progress_callback=profile_progress,
|
||||
graph_id=state.graph_id,
|
||||
parallel_count=parallel_profile_count,
|
||||
realtime_output_path=realtime_output_path,
|
||||
output_platform=realtime_platform
|
||||
)
|
||||
finally:
|
||||
close_reader = getattr(reader, "close", None)
|
||||
if callable(close_reader):
|
||||
close_reader()
|
||||
|
||||
state.profiles_count = len(profiles)
|
||||
|
||||
@@ -448,11 +516,17 @@ class SimulationManager:
|
||||
return state
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"模拟准备失败: {simulation_id}, error={str(e)}")
|
||||
import traceback
|
||||
logger.error(traceback.format_exc())
|
||||
close_reader = locals().get("reader")
|
||||
close_method = getattr(close_reader, "close", None)
|
||||
if callable(close_method):
|
||||
close_method()
|
||||
logger.error(
|
||||
"Simulation preparation failed: simulation_id=%s error_type=%s",
|
||||
simulation_id,
|
||||
type(e).__name__,
|
||||
)
|
||||
state.status = SimulationStatus.FAILED
|
||||
state.error = str(e)
|
||||
state.error = t('api.internalError')
|
||||
self._save_simulation_state(state)
|
||||
raise
|
||||
|
||||
|
||||
@@ -20,8 +20,8 @@ from queue import Queue
|
||||
|
||||
from ..config import Config
|
||||
from ..utils.logger import get_logger
|
||||
from ..utils.locale import get_locale, set_locale
|
||||
from .zep_graph_memory_updater import ZepGraphMemoryManager
|
||||
from ..utils.locale import get_locale, set_locale, t
|
||||
from .local_graph_memory_updater import LocalGraphMemoryManager
|
||||
from .simulation_ipc import SimulationIPCClient, CommandType, IPCResponse
|
||||
|
||||
logger = get_logger('crowdsight.simulation_runner')
|
||||
@@ -308,6 +308,14 @@ class SimulationRunner:
|
||||
json.dump(data, f, ensure_ascii=False, indent=2)
|
||||
|
||||
cls._run_states[state.simulation_id] = state
|
||||
|
||||
@classmethod
|
||||
def _graph_memory_manager(cls):
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
return LocalGraphMemoryManager
|
||||
from .zep_graph_memory_updater import ZepGraphMemoryManager
|
||||
|
||||
return ZepGraphMemoryManager
|
||||
|
||||
@classmethod
|
||||
def start_simulation(
|
||||
@@ -316,7 +324,9 @@ class SimulationRunner:
|
||||
platform: str = "parallel", # twitter / reddit / parallel
|
||||
max_rounds: int = None, # 最大模拟轮数(可选,用于截断过长的模拟)
|
||||
enable_graph_memory_update: bool = False, # 是否将活动更新到Zep图谱
|
||||
graph_id: str = None # Zep图谱ID(启用图谱更新时必需)
|
||||
graph_id: Optional[str] = None, # 图谱ID(启用图谱更新时必需)
|
||||
organization_id: Optional[str] = None,
|
||||
session_factory=None,
|
||||
) -> SimulationRunState:
|
||||
"""
|
||||
启动模拟
|
||||
@@ -375,7 +385,19 @@ class SimulationRunner:
|
||||
raise ValueError("启用图谱记忆更新时必须提供 graph_id")
|
||||
|
||||
try:
|
||||
ZepGraphMemoryManager.create_updater(simulation_id, graph_id)
|
||||
if Config.MEMORY_BACKEND == "local":
|
||||
if not organization_id or session_factory is None:
|
||||
raise ValueError("local_graph_memory_scope_required")
|
||||
LocalGraphMemoryManager.create_updater(
|
||||
simulation_id,
|
||||
graph_id,
|
||||
organization_id=organization_id,
|
||||
session_factory=session_factory,
|
||||
)
|
||||
else:
|
||||
from .zep_graph_memory_updater import ZepGraphMemoryManager
|
||||
|
||||
ZepGraphMemoryManager.create_updater(simulation_id, graph_id)
|
||||
cls._graph_memory_enabled[simulation_id] = True
|
||||
logger.info(f"已启用图谱记忆更新: simulation_id={simulation_id}, graph_id={graph_id}")
|
||||
except Exception as e:
|
||||
@@ -471,8 +493,13 @@ class SimulationRunner:
|
||||
logger.info(f"模拟启动成功: {simulation_id}, pid={process.pid}, platform={platform}")
|
||||
|
||||
except Exception as e:
|
||||
logger.error(
|
||||
"Simulation process start failed: simulation_id=%s error_type=%s",
|
||||
simulation_id,
|
||||
type(e).__name__,
|
||||
)
|
||||
state.runner_status = RunnerStatus.FAILED
|
||||
state.error = str(e)
|
||||
state.error = t('api.internalError')
|
||||
cls._save_run_state(state)
|
||||
raise
|
||||
|
||||
@@ -530,36 +557,38 @@ class SimulationRunner:
|
||||
logger.info(f"模拟完成: {simulation_id}")
|
||||
else:
|
||||
state.runner_status = RunnerStatus.FAILED
|
||||
# 从主日志文件读取错误信息
|
||||
main_log_path = os.path.join(sim_dir, "simulation.log")
|
||||
error_info = ""
|
||||
try:
|
||||
if os.path.exists(main_log_path):
|
||||
with open(main_log_path, 'r', encoding='utf-8') as f:
|
||||
error_info = f.read()[-2000:] # 取最后2000字符
|
||||
except Exception:
|
||||
pass
|
||||
state.error = f"进程退出码: {exit_code}, 错误: {error_info}"
|
||||
logger.error(f"模拟失败: {simulation_id}, error={state.error}")
|
||||
state.error = t('api.internalError')
|
||||
logger.error(
|
||||
"Simulation process failed: simulation_id=%s exit_code=%s",
|
||||
simulation_id,
|
||||
exit_code,
|
||||
)
|
||||
|
||||
state.twitter_running = False
|
||||
state.reddit_running = False
|
||||
cls._save_run_state(state)
|
||||
|
||||
except Exception as e:
|
||||
logger.error(f"监控线程异常: {simulation_id}, error={str(e)}")
|
||||
logger.error(
|
||||
"Simulation monitor failed: simulation_id=%s error_type=%s",
|
||||
simulation_id,
|
||||
type(e).__name__,
|
||||
)
|
||||
state.runner_status = RunnerStatus.FAILED
|
||||
state.error = str(e)
|
||||
state.error = t('api.internalError')
|
||||
cls._save_run_state(state)
|
||||
|
||||
finally:
|
||||
# 停止图谱记忆更新器
|
||||
if cls._graph_memory_enabled.get(simulation_id, False):
|
||||
try:
|
||||
ZepGraphMemoryManager.stop_updater(simulation_id)
|
||||
cls._graph_memory_manager().stop_updater(simulation_id)
|
||||
logger.info(f"已停止图谱记忆更新: simulation_id={simulation_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"停止图谱记忆更新器失败: {e}")
|
||||
logger.error(
|
||||
"Stopping graph memory updater failed: error_type=%s",
|
||||
type(e).__name__,
|
||||
)
|
||||
cls._graph_memory_enabled.pop(simulation_id, None)
|
||||
|
||||
# 清理进程资源
|
||||
@@ -604,7 +633,7 @@ class SimulationRunner:
|
||||
graph_memory_enabled = cls._graph_memory_enabled.get(state.simulation_id, False)
|
||||
graph_updater = None
|
||||
if graph_memory_enabled:
|
||||
graph_updater = ZepGraphMemoryManager.get_updater(state.simulation_id)
|
||||
graph_updater = cls._graph_memory_manager().get_updater(state.simulation_id)
|
||||
|
||||
try:
|
||||
with open(log_path, 'r', encoding='utf-8') as f:
|
||||
@@ -812,7 +841,7 @@ class SimulationRunner:
|
||||
# 停止图谱记忆更新器
|
||||
if cls._graph_memory_enabled.get(simulation_id, False):
|
||||
try:
|
||||
ZepGraphMemoryManager.stop_updater(simulation_id)
|
||||
cls._graph_memory_manager().stop_updater(simulation_id)
|
||||
logger.info(f"已停止图谱记忆更新: simulation_id={simulation_id}")
|
||||
except Exception as e:
|
||||
logger.error(f"停止图谱记忆更新器失败: {e}")
|
||||
@@ -1206,7 +1235,7 @@ class SimulationRunner:
|
||||
|
||||
# 首先停止所有图谱记忆更新器(stop_all 内部会打印日志)
|
||||
try:
|
||||
ZepGraphMemoryManager.stop_all()
|
||||
cls._graph_memory_manager().stop_all()
|
||||
except Exception as e:
|
||||
logger.error(f"停止图谱记忆更新器失败: {e}")
|
||||
cls._graph_memory_enabled.clear()
|
||||
|
||||
69
backend/app/services/usage_service.py
Normal file
69
backend/app/services/usage_service.py
Normal file
@@ -0,0 +1,69 @@
|
||||
"""Durable LLM usage/cost accounting service.
|
||||
|
||||
Records per-organization, per-user LLM usage without storing any prompt content
|
||||
or secrets. A simple default cost estimate (input/output per-token) is applied
|
||||
and can be overridden by a rate table later.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Optional
|
||||
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from ..models.usage import UsageEvent
|
||||
|
||||
# Default per-1K token cost estimates (USD); a rate table can supersede later.
|
||||
_DEFAULT_INPUT_RATE_PER_1K = 0.0025
|
||||
_DEFAULT_OUTPUT_RATE_PER_1K = 0.0100
|
||||
|
||||
|
||||
class UsageService:
|
||||
def __init__(self, session: Session):
|
||||
self.session = session
|
||||
|
||||
def record_event(
|
||||
self,
|
||||
*,
|
||||
organization_id: str,
|
||||
user_id: Optional[str],
|
||||
operation: str,
|
||||
model: Optional[str] = None,
|
||||
input_tokens: int = 0,
|
||||
output_tokens: int = 0,
|
||||
) -> str:
|
||||
if not isinstance(organization_id, str) or not organization_id:
|
||||
raise ValueError("organization_id_required")
|
||||
cost = (
|
||||
(input_tokens / 1000) * _DEFAULT_INPUT_RATE_PER_1K
|
||||
+ (output_tokens / 1000) * _DEFAULT_OUTPUT_RATE_PER_1K
|
||||
)
|
||||
event = UsageEvent(
|
||||
organization_id=organization_id,
|
||||
user_id=user_id,
|
||||
operation=operation,
|
||||
model=model,
|
||||
input_tokens=int(input_tokens or 0),
|
||||
output_tokens=int(output_tokens or 0),
|
||||
estimated_cost=round(cost, 6),
|
||||
)
|
||||
self.session.add(event)
|
||||
self.session.flush()
|
||||
return event.id
|
||||
|
||||
def list_events(self, *, organization_id: str, limit: int = 100) -> list[UsageEvent]:
|
||||
return (
|
||||
self.session.query(UsageEvent)
|
||||
.filter(UsageEvent.organization_id == organization_id)
|
||||
.order_by(UsageEvent.created_at.desc())
|
||||
.limit(min(max(int(limit), 1), 1000))
|
||||
.all()
|
||||
)
|
||||
|
||||
def total_cost(self, *, organization_id: str) -> float:
|
||||
rows = (
|
||||
self.session.query(UsageEvent)
|
||||
.filter(UsageEvent.organization_id == organization_id)
|
||||
.all()
|
||||
)
|
||||
return round(sum(row.estimated_cost for row in rows), 6)
|
||||
@@ -21,182 +21,7 @@ from ..utils.locale import get_locale, set_locale
|
||||
logger = get_logger('crowdsight.zep_graph_memory_updater')
|
||||
|
||||
|
||||
@dataclass
|
||||
class AgentActivity:
|
||||
"""Agent活动记录"""
|
||||
platform: str # twitter / reddit
|
||||
agent_id: int
|
||||
agent_name: str
|
||||
action_type: str # CREATE_POST, LIKE_POST, etc.
|
||||
action_args: Dict[str, Any]
|
||||
round_num: int
|
||||
timestamp: str
|
||||
|
||||
def to_episode_text(self) -> str:
|
||||
"""
|
||||
将活动转换为可以发送给Zep的文本描述
|
||||
|
||||
采用自然语言描述格式,让Zep能够从中提取实体和关系
|
||||
不添加模拟相关的前缀,避免误导图谱更新
|
||||
"""
|
||||
# 根据不同的动作类型生成不同的描述
|
||||
action_descriptions = {
|
||||
"CREATE_POST": self._describe_create_post,
|
||||
"LIKE_POST": self._describe_like_post,
|
||||
"DISLIKE_POST": self._describe_dislike_post,
|
||||
"REPOST": self._describe_repost,
|
||||
"QUOTE_POST": self._describe_quote_post,
|
||||
"FOLLOW": self._describe_follow,
|
||||
"CREATE_COMMENT": self._describe_create_comment,
|
||||
"LIKE_COMMENT": self._describe_like_comment,
|
||||
"DISLIKE_COMMENT": self._describe_dislike_comment,
|
||||
"SEARCH_POSTS": self._describe_search,
|
||||
"SEARCH_USER": self._describe_search_user,
|
||||
"MUTE": self._describe_mute,
|
||||
}
|
||||
|
||||
describe_func = action_descriptions.get(self.action_type, self._describe_generic)
|
||||
description = describe_func()
|
||||
|
||||
# 直接返回 "agent名称: 活动描述" 格式,不添加模拟前缀
|
||||
return f"{self.agent_name}: {description}"
|
||||
|
||||
def _describe_create_post(self) -> str:
|
||||
content = self.action_args.get("content", "")
|
||||
if content:
|
||||
return f"发布了一条帖子:「{content}」"
|
||||
return "发布了一条帖子"
|
||||
|
||||
def _describe_like_post(self) -> str:
|
||||
"""点赞帖子 - 包含帖子原文和作者信息"""
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if post_content and post_author:
|
||||
return f"点赞了{post_author}的帖子:「{post_content}」"
|
||||
elif post_content:
|
||||
return f"点赞了一条帖子:「{post_content}」"
|
||||
elif post_author:
|
||||
return f"点赞了{post_author}的一条帖子"
|
||||
return "点赞了一条帖子"
|
||||
|
||||
def _describe_dislike_post(self) -> str:
|
||||
"""踩帖子 - 包含帖子原文和作者信息"""
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if post_content and post_author:
|
||||
return f"踩了{post_author}的帖子:「{post_content}」"
|
||||
elif post_content:
|
||||
return f"踩了一条帖子:「{post_content}」"
|
||||
elif post_author:
|
||||
return f"踩了{post_author}的一条帖子"
|
||||
return "踩了一条帖子"
|
||||
|
||||
def _describe_repost(self) -> str:
|
||||
"""转发帖子 - 包含原帖内容和作者信息"""
|
||||
original_content = self.action_args.get("original_content", "")
|
||||
original_author = self.action_args.get("original_author_name", "")
|
||||
|
||||
if original_content and original_author:
|
||||
return f"转发了{original_author}的帖子:「{original_content}」"
|
||||
elif original_content:
|
||||
return f"转发了一条帖子:「{original_content}」"
|
||||
elif original_author:
|
||||
return f"转发了{original_author}的一条帖子"
|
||||
return "转发了一条帖子"
|
||||
|
||||
def _describe_quote_post(self) -> str:
|
||||
"""引用帖子 - 包含原帖内容、作者信息和引用评论"""
|
||||
original_content = self.action_args.get("original_content", "")
|
||||
original_author = self.action_args.get("original_author_name", "")
|
||||
quote_content = self.action_args.get("quote_content", "") or self.action_args.get("content", "")
|
||||
|
||||
base = ""
|
||||
if original_content and original_author:
|
||||
base = f"引用了{original_author}的帖子「{original_content}」"
|
||||
elif original_content:
|
||||
base = f"引用了一条帖子「{original_content}」"
|
||||
elif original_author:
|
||||
base = f"引用了{original_author}的一条帖子"
|
||||
else:
|
||||
base = "引用了一条帖子"
|
||||
|
||||
if quote_content:
|
||||
base += f",并评论道:「{quote_content}」"
|
||||
return base
|
||||
|
||||
def _describe_follow(self) -> str:
|
||||
"""关注用户 - 包含被关注用户的名称"""
|
||||
target_user_name = self.action_args.get("target_user_name", "")
|
||||
|
||||
if target_user_name:
|
||||
return f"关注了用户「{target_user_name}」"
|
||||
return "关注了一个用户"
|
||||
|
||||
def _describe_create_comment(self) -> str:
|
||||
"""发表评论 - 包含评论内容和所评论的帖子信息"""
|
||||
content = self.action_args.get("content", "")
|
||||
post_content = self.action_args.get("post_content", "")
|
||||
post_author = self.action_args.get("post_author_name", "")
|
||||
|
||||
if content:
|
||||
if post_content and post_author:
|
||||
return f"在{post_author}的帖子「{post_content}」下评论道:「{content}」"
|
||||
elif post_content:
|
||||
return f"在帖子「{post_content}」下评论道:「{content}」"
|
||||
elif post_author:
|
||||
return f"在{post_author}的帖子下评论道:「{content}」"
|
||||
return f"评论道:「{content}」"
|
||||
return "发表了评论"
|
||||
|
||||
def _describe_like_comment(self) -> str:
|
||||
"""点赞评论 - 包含评论内容和作者信息"""
|
||||
comment_content = self.action_args.get("comment_content", "")
|
||||
comment_author = self.action_args.get("comment_author_name", "")
|
||||
|
||||
if comment_content and comment_author:
|
||||
return f"点赞了{comment_author}的评论:「{comment_content}」"
|
||||
elif comment_content:
|
||||
return f"点赞了一条评论:「{comment_content}」"
|
||||
elif comment_author:
|
||||
return f"点赞了{comment_author}的一条评论"
|
||||
return "点赞了一条评论"
|
||||
|
||||
def _describe_dislike_comment(self) -> str:
|
||||
"""踩评论 - 包含评论内容和作者信息"""
|
||||
comment_content = self.action_args.get("comment_content", "")
|
||||
comment_author = self.action_args.get("comment_author_name", "")
|
||||
|
||||
if comment_content and comment_author:
|
||||
return f"踩了{comment_author}的评论:「{comment_content}」"
|
||||
elif comment_content:
|
||||
return f"踩了一条评论:「{comment_content}」"
|
||||
elif comment_author:
|
||||
return f"踩了{comment_author}的一条评论"
|
||||
return "踩了一条评论"
|
||||
|
||||
def _describe_search(self) -> str:
|
||||
"""搜索帖子 - 包含搜索关键词"""
|
||||
query = self.action_args.get("query", "") or self.action_args.get("keyword", "")
|
||||
return f"搜索了「{query}」" if query else "进行了搜索"
|
||||
|
||||
def _describe_search_user(self) -> str:
|
||||
"""搜索用户 - 包含搜索关键词"""
|
||||
query = self.action_args.get("query", "") or self.action_args.get("username", "")
|
||||
return f"搜索了用户「{query}」" if query else "搜索了用户"
|
||||
|
||||
def _describe_mute(self) -> str:
|
||||
"""屏蔽用户 - 包含被屏蔽用户的名称"""
|
||||
target_user_name = self.action_args.get("target_user_name", "")
|
||||
|
||||
if target_user_name:
|
||||
return f"屏蔽了用户「{target_user_name}」"
|
||||
return "屏蔽了一个用户"
|
||||
|
||||
def _describe_generic(self) -> str:
|
||||
# 对于未知的动作类型,生成通用描述
|
||||
return f"执行了{self.action_type}操作"
|
||||
from .memory_activity import AgentActivity
|
||||
|
||||
|
||||
class ZepGraphMemoryUpdater:
|
||||
|
||||
@@ -1461,14 +1461,18 @@ Return the sub-questions in JSON format."""
|
||||
|
||||
except ValueError as e:
|
||||
# 模拟环境未运行
|
||||
logger.warning(t("console.interviewApiCallFailed", error=e))
|
||||
result.summary = f"Interview failed: {str(e)}. Simulation environment may have closed. Please ensure OASIS environment is running."
|
||||
logger.warning(
|
||||
"Interview API call failed: error_type=%s",
|
||||
type(e).__name__,
|
||||
)
|
||||
result.summary = t("api.internalError")
|
||||
return result
|
||||
except Exception as e:
|
||||
logger.error(t("console.interviewApiCallException", error=e))
|
||||
import traceback
|
||||
logger.error(traceback.format_exc())
|
||||
result.summary = f"采访过程发生错误:{str(e)}"
|
||||
logger.error(
|
||||
"Interview API call exception: error_type=%s",
|
||||
type(e).__name__,
|
||||
)
|
||||
result.summary = t("api.internalError")
|
||||
return result
|
||||
|
||||
# Step 6: 生成采访摘要
|
||||
|
||||
43
backend/app/utils/api_errors.py
Normal file
43
backend/app/utils/api_errors.py
Normal file
@@ -0,0 +1,43 @@
|
||||
"""Structured, localized API error contracts.
|
||||
|
||||
The exception text is intentionally never serialized. Route handlers can raise
|
||||
``ApiError`` with a stable code and translation key; Flask integration resolves
|
||||
the message through the current locale at the response boundary.
|
||||
"""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from dataclasses import dataclass, field
|
||||
from typing import Callable, Mapping
|
||||
|
||||
|
||||
Translator = Callable[..., str]
|
||||
|
||||
|
||||
@dataclass
|
||||
class ApiError(Exception):
|
||||
"""A safe application error that can cross the HTTP boundary."""
|
||||
|
||||
code: str
|
||||
status_code: int
|
||||
message_key: str
|
||||
params: Mapping[str, object] = field(default_factory=dict)
|
||||
|
||||
def __post_init__(self):
|
||||
Exception.__init__(self, self.code)
|
||||
|
||||
def to_payload(self, translate: Translator) -> dict[str, object]:
|
||||
return {
|
||||
"success": False,
|
||||
"error_code": self.code,
|
||||
"message": translate(self.message_key, **dict(self.params)),
|
||||
}
|
||||
|
||||
|
||||
def internal_error_payload(translate: Translator) -> dict[str, object]:
|
||||
"""Return a generic internal error without accepting exception details."""
|
||||
return ApiError(
|
||||
code="internal_error",
|
||||
status_code=500,
|
||||
message_key="api.internalError",
|
||||
).to_payload(translate)
|
||||
82
backend/app/utils/language_policy.py
Normal file
82
backend/app/utils/language_policy.py
Normal file
@@ -0,0 +1,82 @@
|
||||
"""Language policy shared by API locale negotiation and background jobs."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
from typing import Iterable
|
||||
|
||||
SUPPORTED_LOCALES = ("th", "en")
|
||||
DEFAULT_LOCALE = "th"
|
||||
|
||||
# Legacy values are intentionally mapped to the default locale so an old
|
||||
# browser preference cannot re-enable an unsupported product language.
|
||||
_LEGACY_LOCALE_ALIASES = {
|
||||
"zh": DEFAULT_LOCALE,
|
||||
"zh-cn": DEFAULT_LOCALE,
|
||||
"zh-tw": DEFAULT_LOCALE,
|
||||
}
|
||||
|
||||
|
||||
def normalize_locale(value: object, default: str = DEFAULT_LOCALE) -> str:
|
||||
"""Return one supported base locale, failing closed to ``default``."""
|
||||
safe_default = default if default in SUPPORTED_LOCALES else DEFAULT_LOCALE
|
||||
if not isinstance(value, str):
|
||||
return safe_default
|
||||
|
||||
normalized = value.strip().replace("_", "-").lower()
|
||||
if not normalized:
|
||||
return safe_default
|
||||
if normalized in _LEGACY_LOCALE_ALIASES:
|
||||
return _LEGACY_LOCALE_ALIASES[normalized]
|
||||
if normalized in SUPPORTED_LOCALES:
|
||||
return normalized
|
||||
|
||||
base_locale = normalized.split("-", 1)[0]
|
||||
if base_locale in SUPPORTED_LOCALES:
|
||||
return base_locale
|
||||
if base_locale in _LEGACY_LOCALE_ALIASES:
|
||||
return _LEGACY_LOCALE_ALIASES[base_locale]
|
||||
return safe_default
|
||||
|
||||
|
||||
def locale_from_accept_language(header: object) -> str:
|
||||
"""Choose the best supported locale from an HTTP Accept-Language value."""
|
||||
if not isinstance(header, str) or not header.strip():
|
||||
return DEFAULT_LOCALE
|
||||
|
||||
candidates: list[tuple[float, int, str]] = []
|
||||
for position, raw_item in enumerate(header.split(",")):
|
||||
parts = [part.strip() for part in raw_item.split(";")]
|
||||
language = parts[0]
|
||||
quality = 1.0
|
||||
for parameter in parts[1:]:
|
||||
key, separator, value = parameter.partition("=")
|
||||
if key.strip().lower() != "q" or not separator:
|
||||
continue
|
||||
try:
|
||||
quality = float(value.strip())
|
||||
except ValueError:
|
||||
quality = 0.0
|
||||
break
|
||||
|
||||
if quality <= 0:
|
||||
continue
|
||||
normalized_language = language.strip().replace("_", "-").lower()
|
||||
base_locale = normalized_language.split("-", 1)[0]
|
||||
if base_locale in SUPPORTED_LOCALES:
|
||||
# Higher q wins; original order breaks ties.
|
||||
candidates.append((quality, -position, base_locale))
|
||||
|
||||
if not candidates:
|
||||
return DEFAULT_LOCALE
|
||||
candidates.sort(reverse=True)
|
||||
return candidates[0][2]
|
||||
|
||||
|
||||
def is_supported_locale(value: object) -> bool:
|
||||
"""Return whether ``value`` is already a canonical supported locale."""
|
||||
return isinstance(value, str) and value in SUPPORTED_LOCALES
|
||||
|
||||
|
||||
def supported_locales() -> Iterable[str]:
|
||||
"""Expose supported locales without allowing callers to mutate the tuple."""
|
||||
return SUPPORTED_LOCALES
|
||||
@@ -3,6 +3,13 @@ import os
|
||||
import threading
|
||||
from flask import request, has_request_context
|
||||
|
||||
from .language_policy import (
|
||||
DEFAULT_LOCALE,
|
||||
SUPPORTED_LOCALES,
|
||||
locale_from_accept_language,
|
||||
normalize_locale,
|
||||
)
|
||||
|
||||
_thread_local = threading.local()
|
||||
|
||||
_locales_dir = os.path.join(os.path.dirname(__file__), '..', '..', '..', 'locales')
|
||||
@@ -16,25 +23,26 @@ _translations = {}
|
||||
for filename in os.listdir(_locales_dir):
|
||||
if filename.endswith('.json') and filename != 'languages.json':
|
||||
locale_name = filename[:-5]
|
||||
if locale_name not in SUPPORTED_LOCALES:
|
||||
continue
|
||||
with open(os.path.join(_locales_dir, filename), 'r', encoding='utf-8') as f:
|
||||
_translations[locale_name] = json.load(f)
|
||||
|
||||
|
||||
def set_locale(locale: str):
|
||||
"""Set locale for current thread. Call at the start of background threads."""
|
||||
_thread_local.locale = locale
|
||||
"""Set a canonical locale for the current background thread."""
|
||||
_thread_local.locale = normalize_locale(locale)
|
||||
|
||||
|
||||
def get_locale() -> str:
|
||||
if has_request_context():
|
||||
raw = request.headers.get('Accept-Language', 'zh')
|
||||
return raw if raw in _translations else 'zh'
|
||||
return getattr(_thread_local, 'locale', 'zh')
|
||||
return locale_from_accept_language(request.headers.get('Accept-Language', ''))
|
||||
return normalize_locale(getattr(_thread_local, 'locale', DEFAULT_LOCALE))
|
||||
|
||||
|
||||
def t(key: str, **kwargs) -> str:
|
||||
locale = get_locale()
|
||||
messages = _translations.get(locale, _translations.get('zh', {}))
|
||||
messages = _translations.get(locale, _translations.get(DEFAULT_LOCALE, {}))
|
||||
|
||||
value = messages
|
||||
for part in key.split('.'):
|
||||
@@ -45,7 +53,7 @@ def t(key: str, **kwargs) -> str:
|
||||
break
|
||||
|
||||
if value is None:
|
||||
value = _translations.get('zh', {})
|
||||
value = _translations.get(DEFAULT_LOCALE, {})
|
||||
for part in key.split('.'):
|
||||
if isinstance(value, dict):
|
||||
value = value.get(part)
|
||||
@@ -65,5 +73,10 @@ def t(key: str, **kwargs) -> str:
|
||||
|
||||
def get_language_instruction() -> str:
|
||||
locale = get_locale()
|
||||
lang_config = _languages.get(locale, _languages.get('zh', {}))
|
||||
return lang_config.get('llmInstruction', '请使用中文回答。')
|
||||
lang_config = _languages.get(locale, _languages.get(DEFAULT_LOCALE, {}))
|
||||
return lang_config.get(
|
||||
'llmInstruction',
|
||||
'IMPORTANT: Respond exclusively in Thai language.'
|
||||
if locale == DEFAULT_LOCALE
|
||||
else 'IMPORTANT: Respond exclusively in English language.',
|
||||
)
|
||||
|
||||
58
backend/migrations/env.py
Normal file
58
backend/migrations/env.py
Normal file
@@ -0,0 +1,58 @@
|
||||
"""Alembic environment for the SaaS identity schema."""
|
||||
|
||||
from __future__ import annotations
|
||||
|
||||
import os
|
||||
from logging.config import fileConfig
|
||||
|
||||
from alembic import context
|
||||
from sqlalchemy import pool
|
||||
|
||||
from app.db import Base, create_database_engine
|
||||
from app.models import memory, saas # noqa: F401 - register models on metadata
|
||||
|
||||
config = context.config
|
||||
if config.config_file_name is not None:
|
||||
fileConfig(config.config_file_name)
|
||||
|
||||
|
||||
target_metadata = Base.metadata
|
||||
|
||||
|
||||
def database_url() -> str:
|
||||
url = os.environ.get("DATABASE_URL") or config.get_main_option("sqlalchemy.url")
|
||||
if not url:
|
||||
raise RuntimeError("DATABASE_URL is required for migrations")
|
||||
return url
|
||||
|
||||
|
||||
def run_migrations_offline() -> None:
|
||||
context.configure(
|
||||
url=database_url(),
|
||||
target_metadata=target_metadata,
|
||||
literal_binds=True,
|
||||
dialect_opts={"paramstyle": "named"},
|
||||
compare_type=True,
|
||||
compare_server_default=True,
|
||||
)
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
def run_migrations_online() -> None:
|
||||
connectable = create_database_engine(database_url(), poolclass=pool.NullPool)
|
||||
with connectable.connect() as connection:
|
||||
context.configure(
|
||||
connection=connection,
|
||||
target_metadata=target_metadata,
|
||||
compare_type=True,
|
||||
compare_server_default=True,
|
||||
)
|
||||
with context.begin_transaction():
|
||||
context.run_migrations()
|
||||
|
||||
|
||||
if context.is_offline_mode():
|
||||
run_migrations_offline()
|
||||
else:
|
||||
run_migrations_online()
|
||||
26
backend/migrations/script.py.mako
Normal file
26
backend/migrations/script.py.mako
Normal file
@@ -0,0 +1,26 @@
|
||||
"""${message}
|
||||
|
||||
Revision ID: ${up_revision}
|
||||
Revises: ${down_revision | comma,n}
|
||||
Create Date: ${create_date}
|
||||
"""
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
${imports if imports else ""}
|
||||
|
||||
# revision identifiers, used by Alembic.
|
||||
revision: str = ${repr(up_revision)}
|
||||
down_revision: Union[str, None] = ${repr(down_revision)}
|
||||
branch_labels: Union[str, Sequence[str], None] = ${repr(branch_labels)}
|
||||
depends_on: Union[str, Sequence[str], None] = ${repr(depends_on)}
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
${upgrades if upgrades else "pass"}
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
${downgrades if downgrades else "pass"}
|
||||
79
backend/migrations/versions/0001_identity.py
Normal file
79
backend/migrations/versions/0001_identity.py
Normal file
@@ -0,0 +1,79 @@
|
||||
"""Create SaaS identity and tenant membership tables.
|
||||
|
||||
Revision ID: 0001_identity
|
||||
Revises:
|
||||
Create Date: 2026-08-23
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "0001_identity"
|
||||
down_revision: Union[str, None] = None
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"organizations",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("name", sa.String(length=160), nullable=False),
|
||||
sa.Column("slug", sa.String(length=80), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'active'"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_organizations_slug", "organizations", ["slug"], unique=True)
|
||||
|
||||
op.create_table(
|
||||
"users",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("email_normalized", sa.String(length=320), nullable=False),
|
||||
sa.Column(
|
||||
"password_hash",
|
||||
sa.String(length=512),
|
||||
server_default=sa.text("'!invite_pending'"),
|
||||
nullable=False,
|
||||
),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'active'"), nullable=False),
|
||||
sa.Column("auth_version", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("locale", sa.String(length=8), server_default=sa.text("'th'"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("last_login_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_users_email_normalized", "users", ["email_normalized"], unique=True)
|
||||
|
||||
op.create_table(
|
||||
"memberships",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("role", sa.String(length=11), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'active'"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.CheckConstraint(
|
||||
"role IN ('super_admin', 'admin', 'user')",
|
||||
name="ck_membership_role",
|
||||
),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("user_id", "organization_id", name="uq_membership_user_org"),
|
||||
)
|
||||
op.create_index("ix_memberships_user_id", "memberships", ["user_id"], unique=False)
|
||||
op.create_index("ix_memberships_organization_id", "memberships", ["organization_id"], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_memberships_organization_id", table_name="memberships")
|
||||
op.drop_index("ix_memberships_user_id", table_name="memberships")
|
||||
op.drop_table("memberships")
|
||||
op.drop_index("ix_users_email_normalized", table_name="users")
|
||||
op.drop_table("users")
|
||||
op.drop_index("ix_organizations_slug", table_name="organizations")
|
||||
op.drop_table("organizations")
|
||||
45
backend/migrations/versions/0002_sessions.py
Normal file
45
backend/migrations/versions/0002_sessions.py
Normal file
@@ -0,0 +1,45 @@
|
||||
"""Add revocable opaque authentication sessions.
|
||||
|
||||
Revision ID: 0002_sessions
|
||||
Revises: 0001_identity
|
||||
Create Date: 2026-08-23
|
||||
"""
|
||||
|
||||
from typing import Sequence, Union
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision: str = "0002_sessions"
|
||||
down_revision: Union[str, None] = "0001_identity"
|
||||
branch_labels: Union[str, Sequence[str], None] = None
|
||||
depends_on: Union[str, Sequence[str], None] = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"sessions",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("membership_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("token_hash", sa.String(length=64), nullable=False),
|
||||
sa.Column("auth_version", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("revoked_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("last_seen_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.ForeignKeyConstraint(["membership_id"], ["memberships.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_sessions_user_id", "sessions", ["user_id"], unique=False)
|
||||
op.create_index("ix_sessions_membership_id", "sessions", ["membership_id"], unique=False)
|
||||
op.create_index("ix_sessions_token_hash", "sessions", ["token_hash"], unique=True)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_sessions_token_hash", table_name="sessions")
|
||||
op.drop_index("ix_sessions_membership_id", table_name="sessions")
|
||||
op.drop_index("ix_sessions_user_id", table_name="sessions")
|
||||
op.drop_table("sessions")
|
||||
116
backend/migrations/versions/0003_memory.py
Normal file
116
backend/migrations/versions/0003_memory.py
Normal file
@@ -0,0 +1,116 @@
|
||||
"""Add durable local graph-memory tables."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision = "0003_memory"
|
||||
down_revision = "0002_sessions"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"memory_graphs",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("project_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("ontology", sa.JSON(), server_default=sa.text("'{}'"), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'ready'"), nullable=False),
|
||||
sa.Column("version", sa.Integer(), server_default=sa.text("1"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_memory_graphs_organization_id", "memory_graphs", ["organization_id"], unique=False)
|
||||
op.create_index("ix_memory_graphs_project_id", "memory_graphs", ["project_id"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"memory_episodes",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("graph_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("source_type", sa.String(length=32), nullable=False),
|
||||
sa.Column("source_ref", sa.String(length=256), nullable=False),
|
||||
sa.Column("normalized_text", sa.Text(), nullable=False),
|
||||
sa.Column("summary", sa.Text(), server_default=sa.text("''"), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'processed'"), nullable=False),
|
||||
sa.Column("extractor_version", sa.String(length=64), server_default=sa.text("'v1'"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["graph_id"], ["memory_graphs.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("graph_id", "source_type", "source_ref", name="uq_memory_episode_source"),
|
||||
)
|
||||
op.create_index("ix_memory_episodes_graph_id", "memory_episodes", ["graph_id"], unique=False)
|
||||
op.create_index("ix_memory_episodes_graph_status", "memory_episodes", ["graph_id", "status"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"memory_nodes",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("graph_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("canonical_name", sa.String(length=512), nullable=False),
|
||||
sa.Column("normalized_name", sa.String(length=512), nullable=False),
|
||||
sa.Column("labels", sa.JSON(), server_default=sa.text("'[]'"), nullable=False),
|
||||
sa.Column("aliases", sa.JSON(), server_default=sa.text("'[]'"), nullable=False),
|
||||
sa.Column("attributes", sa.JSON(), server_default=sa.text("'{}'"), nullable=False),
|
||||
sa.Column("summary", sa.Text(), server_default=sa.text("''"), nullable=False),
|
||||
sa.Column("confidence", sa.Float(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["graph_id"], ["memory_graphs.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("graph_id", "normalized_name", name="uq_memory_node_graph_name"),
|
||||
)
|
||||
op.create_index("ix_memory_nodes_graph_id", "memory_nodes", ["graph_id"], unique=False)
|
||||
op.create_index("ix_memory_nodes_graph_name", "memory_nodes", ["graph_id", "normalized_name"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"memory_edges",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("graph_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("source_node_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("target_node_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("relation", sa.String(length=128), nullable=False),
|
||||
sa.Column("fact", sa.Text(), nullable=False),
|
||||
sa.Column("attributes", sa.JSON(), server_default=sa.text("'{}'"), nullable=False),
|
||||
sa.Column("confidence", sa.Float(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("valid_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("invalid_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("expired_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["graph_id"], ["memory_graphs.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["source_node_id"], ["memory_nodes.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["target_node_id"], ["memory_nodes.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_memory_edges_graph_id", "memory_edges", ["graph_id"], unique=False)
|
||||
op.create_index("ix_memory_edges_source_node_id", "memory_edges", ["source_node_id"], unique=False)
|
||||
op.create_index("ix_memory_edges_target_node_id", "memory_edges", ["target_node_id"], unique=False)
|
||||
op.create_index("ix_memory_edges_graph_relation", "memory_edges", ["graph_id", "relation"], unique=False)
|
||||
op.create_index(
|
||||
"ix_memory_edges_graph_temporal",
|
||||
"memory_edges",
|
||||
["graph_id", "valid_at", "invalid_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_memory_edges_graph_temporal", table_name="memory_edges")
|
||||
op.drop_index("ix_memory_edges_graph_relation", table_name="memory_edges")
|
||||
op.drop_index("ix_memory_edges_target_node_id", table_name="memory_edges")
|
||||
op.drop_index("ix_memory_edges_source_node_id", table_name="memory_edges")
|
||||
op.drop_index("ix_memory_edges_graph_id", table_name="memory_edges")
|
||||
op.drop_table("memory_edges")
|
||||
|
||||
op.drop_index("ix_memory_nodes_graph_name", table_name="memory_nodes")
|
||||
op.drop_index("ix_memory_nodes_graph_id", table_name="memory_nodes")
|
||||
op.drop_table("memory_nodes")
|
||||
|
||||
op.drop_index("ix_memory_episodes_graph_status", table_name="memory_episodes")
|
||||
op.drop_index("ix_memory_episodes_graph_id", table_name="memory_episodes")
|
||||
op.drop_table("memory_episodes")
|
||||
|
||||
op.drop_index("ix_memory_graphs_project_id", table_name="memory_graphs")
|
||||
op.drop_index("ix_memory_graphs_organization_id", table_name="memory_graphs")
|
||||
op.drop_table("memory_graphs")
|
||||
107
backend/migrations/versions/0004_operations.py
Normal file
107
backend/migrations/versions/0004_operations.py
Normal file
@@ -0,0 +1,107 @@
|
||||
"""Add durable jobs, idempotency records, and audit logs."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision = "0004_operations"
|
||||
down_revision = "0003_memory"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"jobs",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("owner_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("project_id", sa.String(length=128), nullable=True),
|
||||
sa.Column("graph_id", sa.String(length=128), nullable=True),
|
||||
sa.Column("operation", sa.String(length=120), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'queued'"), nullable=False),
|
||||
sa.Column("progress", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("error_code", sa.String(length=120), nullable=True),
|
||||
sa.Column("result_ref", sa.String(length=512), nullable=True),
|
||||
sa.Column("idempotency_key", sa.String(length=128), nullable=True),
|
||||
sa.Column("attempt", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("settings_version", sa.String(length=128), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("finished_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.CheckConstraint(
|
||||
"status IN ('queued', 'running', 'succeeded', 'failed', 'cancelled')",
|
||||
name="ck_jobs_status",
|
||||
),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["owner_user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_jobs_organization_id", "jobs", ["organization_id"], unique=False)
|
||||
op.create_index("ix_jobs_owner_user_id", "jobs", ["owner_user_id"], unique=False)
|
||||
op.create_index("ix_jobs_project_id", "jobs", ["project_id"], unique=False)
|
||||
op.create_index("ix_jobs_graph_id", "jobs", ["graph_id"], unique=False)
|
||||
op.create_index("ix_jobs_org_status_created", "jobs", ["organization_id", "status", "created_at"], unique=False)
|
||||
op.create_index("ix_jobs_org_owner", "jobs", ["organization_id", "owner_user_id"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"idempotency_records",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("key", sa.String(length=128), nullable=False),
|
||||
sa.Column("request_hash", sa.String(length=64), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'reserved'"), nullable=False),
|
||||
sa.Column("response_status", sa.Integer(), nullable=True),
|
||||
sa.Column("response_body", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("completed_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("organization_id", "user_id", "key", name="uq_idempotency_org_user_key"),
|
||||
)
|
||||
op.create_index("ix_idempotency_records_organization_id", "idempotency_records", ["organization_id"], unique=False)
|
||||
op.create_index("ix_idempotency_records_user_id", "idempotency_records", ["user_id"], unique=False)
|
||||
op.create_index("ix_idempotency_expiry", "idempotency_records", ["expires_at"], unique=False)
|
||||
|
||||
op.create_table(
|
||||
"audit_logs",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("actor_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("action", sa.String(length=160), nullable=False),
|
||||
sa.Column("target_type", sa.String(length=80), nullable=False),
|
||||
sa.Column("target_id", sa.String(length=160), nullable=True),
|
||||
sa.Column("metadata", sa.JSON(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["actor_user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_audit_logs_organization_id", "audit_logs", ["organization_id"], unique=False)
|
||||
op.create_index("ix_audit_logs_actor_user_id", "audit_logs", ["actor_user_id"], unique=False)
|
||||
op.create_index("ix_audit_org_created", "audit_logs", ["organization_id", "created_at"], unique=False)
|
||||
op.create_index("ix_audit_org_target", "audit_logs", ["organization_id", "target_type", "target_id"], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_audit_org_target", table_name="audit_logs")
|
||||
op.drop_index("ix_audit_org_created", table_name="audit_logs")
|
||||
op.drop_index("ix_audit_logs_actor_user_id", table_name="audit_logs")
|
||||
op.drop_index("ix_audit_logs_organization_id", table_name="audit_logs")
|
||||
op.drop_table("audit_logs")
|
||||
|
||||
op.drop_index("ix_idempotency_expiry", table_name="idempotency_records")
|
||||
op.drop_index("ix_idempotency_records_user_id", table_name="idempotency_records")
|
||||
op.drop_index("ix_idempotency_records_organization_id", table_name="idempotency_records")
|
||||
op.drop_table("idempotency_records")
|
||||
|
||||
op.drop_index("ix_jobs_org_owner", table_name="jobs")
|
||||
op.drop_index("ix_jobs_org_status_created", table_name="jobs")
|
||||
op.drop_index("ix_jobs_graph_id", table_name="jobs")
|
||||
op.drop_index("ix_jobs_project_id", table_name="jobs")
|
||||
op.drop_index("ix_jobs_owner_user_id", table_name="jobs")
|
||||
op.drop_index("ix_jobs_organization_id", table_name="jobs")
|
||||
op.drop_table("jobs")
|
||||
22
backend/migrations/versions/0005_job_payload.py
Normal file
22
backend/migrations/versions/0005_job_payload.py
Normal file
@@ -0,0 +1,22 @@
|
||||
"""Persist task message, result, and progress detail in durable jobs."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision = "0005_job_payload"
|
||||
down_revision = "0004_operations"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("jobs", sa.Column("message", sa.Text(), server_default=sa.text("''"), nullable=False))
|
||||
op.add_column("jobs", sa.Column("result", sa.JSON(), nullable=True))
|
||||
op.add_column("jobs", sa.Column("progress_detail", sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("jobs", "progress_detail")
|
||||
op.drop_column("jobs", "result")
|
||||
op.drop_column("jobs", "message")
|
||||
18
backend/migrations/versions/0006_job_metadata.py
Normal file
18
backend/migrations/versions/0006_job_metadata.py
Normal file
@@ -0,0 +1,18 @@
|
||||
"""Add durable job metadata payload."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
|
||||
revision = "0006_job_metadata"
|
||||
down_revision = "0005_job_payload"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.add_column("jobs", sa.Column("metadata", sa.JSON(), nullable=True))
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_column("jobs", "metadata")
|
||||
119
backend/migrations/versions/0007_product_resources.py
Normal file
119
backend/migrations/versions/0007_product_resources.py
Normal file
@@ -0,0 +1,119 @@
|
||||
"""Add durable product-resource tables: projects, simulations, and reports.
|
||||
|
||||
These replace the legacy filesystem-backed project/simulation/report payloads
|
||||
with tenant- and owner-scoped SQL rows.
|
||||
"""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0007_product_resources"
|
||||
down_revision = "0006_job_metadata"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"projects",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("owner_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("name", sa.String(length=255), server_default=sa.text("''"), nullable=False),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'created'"), nullable=False),
|
||||
sa.Column("language", sa.String(length=16), server_default=sa.text("'en'"), nullable=False),
|
||||
sa.Column("total_text_length", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("source_metadata", sa.JSON(), nullable=True),
|
||||
sa.Column("ontology", sa.JSON(), nullable=True),
|
||||
sa.Column("analysis_summary", sa.Text(), nullable=True),
|
||||
sa.Column("simulation_requirement", sa.Text(), nullable=True),
|
||||
sa.Column("graph_id", sa.String(length=128), nullable=True),
|
||||
sa.Column("graph_build_task_id", sa.String(length=128), nullable=True),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["owner_user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_projects_organization_id", "projects", ["organization_id"])
|
||||
op.create_index("ix_projects_owner_user_id", "projects", ["owner_user_id"])
|
||||
op.create_index("ix_projects_graph_id", "projects", ["graph_id"])
|
||||
op.create_index("ix_projects_org_owner", "projects", ["organization_id", "owner_user_id"])
|
||||
op.create_index("ix_projects_org_created", "projects", ["organization_id", "created_at"])
|
||||
|
||||
op.create_table(
|
||||
"simulations",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("project_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("created_by_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'created'"), nullable=False),
|
||||
sa.Column("platform", sa.String(length=32), server_default=sa.text("'parallel'"), nullable=False),
|
||||
sa.Column("config", sa.JSON(), nullable=True),
|
||||
sa.Column("current_round", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("finished_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["project_id"], ["projects.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["created_by_user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_simulations_organization_id", "simulations", ["organization_id"])
|
||||
op.create_index("ix_simulations_project_id", "simulations", ["project_id"])
|
||||
op.create_index("ix_simulations_org_project", "simulations", ["organization_id", "project_id"])
|
||||
op.create_index("ix_simulations_org_created", "simulations", ["organization_id", "created_at"])
|
||||
|
||||
op.create_table(
|
||||
"reports",
|
||||
sa.Column("id", sa.String(length=128), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("project_id", sa.String(length=128), nullable=False),
|
||||
sa.Column("simulation_id", sa.String(length=128), nullable=True),
|
||||
sa.Column("created_by_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("status", sa.String(length=32), server_default=sa.text("'draft'"), nullable=False),
|
||||
sa.Column("title", sa.String(length=255), server_default=sa.text("''"), nullable=False),
|
||||
sa.Column("outline", sa.JSON(), nullable=True),
|
||||
sa.Column("markdown_content", sa.Text(), nullable=True),
|
||||
sa.Column("error", sa.Text(), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("updated_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.Column("finished_at", sa.DateTime(timezone=True), nullable=True),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["project_id"], ["projects.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["simulation_id"], ["simulations.id"], ondelete="SET NULL"),
|
||||
sa.ForeignKeyConstraint(["created_by_user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("organization_id", "id", name="uq_reports_org_id"),
|
||||
)
|
||||
op.create_index("ix_reports_organization_id", "reports", ["organization_id"])
|
||||
op.create_index("ix_reports_project_id", "reports", ["project_id"])
|
||||
op.create_index("ix_reports_simulation_id", "reports", ["simulation_id"])
|
||||
op.create_index("ix_reports_org_project", "reports", ["organization_id", "project_id"])
|
||||
op.create_index("ix_reports_org_simulation", "reports", ["organization_id", "simulation_id"])
|
||||
op.create_index("ix_reports_org_created", "reports", ["organization_id", "created_at"])
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_reports_org_created", table_name="reports")
|
||||
op.drop_index("ix_reports_org_simulation", table_name="reports")
|
||||
op.drop_index("ix_reports_org_project", table_name="reports")
|
||||
op.drop_index("ix_reports_simulation_id", table_name="reports")
|
||||
op.drop_index("ix_reports_project_id", table_name="reports")
|
||||
op.drop_index("ix_reports_organization_id", table_name="reports")
|
||||
op.drop_table("reports")
|
||||
|
||||
op.drop_index("ix_simulations_org_created", table_name="simulations")
|
||||
op.drop_index("ix_simulations_org_project", table_name="simulations")
|
||||
op.drop_index("ix_simulations_project_id", table_name="simulations")
|
||||
op.drop_index("ix_simulations_organization_id", table_name="simulations")
|
||||
op.drop_table("simulations")
|
||||
|
||||
op.drop_index("ix_projects_org_created", table_name="projects")
|
||||
op.drop_index("ix_projects_org_owner", table_name="projects")
|
||||
op.drop_index("ix_projects_graph_id", table_name="projects")
|
||||
op.drop_index("ix_projects_owner_user_id", table_name="projects")
|
||||
op.drop_index("ix_projects_organization_id", table_name="projects")
|
||||
op.drop_table("projects")
|
||||
28
backend/migrations/versions/0008_platform_settings.py
Normal file
28
backend/migrations/versions/0008_platform_settings.py
Normal file
@@ -0,0 +1,28 @@
|
||||
"""Add durable, versioned platform settings with redacted secrets."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0008_platform_settings"
|
||||
down_revision = "0007_product_resources"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"platform_settings",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("version", sa.String(length=64), nullable=False),
|
||||
sa.Column("settings", sa.JSON(), nullable=True),
|
||||
sa.Column("secret_ref", sa.String(length=512), nullable=True),
|
||||
sa.Column("updated_by_user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("active", sa.Boolean(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
sa.UniqueConstraint("version", name="uq_platform_settings_version"),
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_table("platform_settings")
|
||||
41
backend/migrations/versions/0009_rate_limit.py
Normal file
41
backend/migrations/versions/0009_rate_limit.py
Normal file
@@ -0,0 +1,41 @@
|
||||
"""Add durable rate-limit event records."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0009_rate_limit"
|
||||
down_revision = "0008_platform_settings"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"rate_limit_events",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("operation", sa.String(length=120), nullable=False),
|
||||
sa.Column("key", sa.String(length=255), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_rate_limit_events_key", "rate_limit_events", ["key"], unique=False)
|
||||
op.create_index(
|
||||
"ix_rate_limit_op_key_created",
|
||||
"rate_limit_events",
|
||||
["operation", "key", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
op.create_index(
|
||||
"ix_rate_limit_org_created",
|
||||
"rate_limit_events",
|
||||
["organization_id", "created_at"],
|
||||
unique=False,
|
||||
)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_rate_limit_org_created", table_name="rate_limit_events")
|
||||
op.drop_index("ix_rate_limit_op_key_created", table_name="rate_limit_events")
|
||||
op.drop_index("ix_rate_limit_events_key", table_name="rate_limit_events")
|
||||
op.drop_table("rate_limit_events")
|
||||
39
backend/migrations/versions/0010_usage_events.py
Normal file
39
backend/migrations/versions/0010_usage_events.py
Normal file
@@ -0,0 +1,39 @@
|
||||
"""Add durable LLM usage/cost events."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0010_usage_events"
|
||||
down_revision = "0009_rate_limit"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"usage_events",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("organization_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=64), nullable=True),
|
||||
sa.Column("operation", sa.String(length=160), nullable=False),
|
||||
sa.Column("model", sa.String(length=120), nullable=True),
|
||||
sa.Column("input_tokens", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("output_tokens", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("estimated_cost", sa.Float(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["organization_id"], ["organizations.id"], ondelete="CASCADE"),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="SET NULL"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_usage_events_organization_id", "usage_events", ["organization_id"], unique=False)
|
||||
op.create_index("ix_usage_events_user_id", "usage_events", ["user_id"], unique=False)
|
||||
op.create_index("ix_usage_org_created", "usage_events", ["organization_id", "created_at"], unique=False)
|
||||
op.create_index("ix_usage_org_user", "usage_events", ["organization_id", "user_id"], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_usage_org_user", table_name="usage_events")
|
||||
op.drop_index("ix_usage_org_created", table_name="usage_events")
|
||||
op.drop_index("ix_usage_events_user_id", table_name="usage_events")
|
||||
op.drop_index("ix_usage_events_organization_id", table_name="usage_events")
|
||||
op.drop_table("usage_events")
|
||||
30
backend/migrations/versions/0011_password_reset_tokens.py
Normal file
30
backend/migrations/versions/0011_password_reset_tokens.py
Normal file
@@ -0,0 +1,30 @@
|
||||
"""Add durable, single-use password reset tokens."""
|
||||
|
||||
from alembic import op
|
||||
import sqlalchemy as sa
|
||||
|
||||
revision = "0011_password_reset_tokens"
|
||||
down_revision = "0010_usage_events"
|
||||
branch_labels = None
|
||||
depends_on = None
|
||||
|
||||
|
||||
def upgrade() -> None:
|
||||
op.create_table(
|
||||
"password_reset_tokens",
|
||||
sa.Column("id", sa.String(length=64), nullable=False),
|
||||
sa.Column("user_id", sa.String(length=64), nullable=False),
|
||||
sa.Column("token_hash", sa.String(length=128), nullable=False),
|
||||
sa.Column("auth_version", sa.Integer(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("used", sa.Boolean(), server_default=sa.text("0"), nullable=False),
|
||||
sa.Column("expires_at", sa.DateTime(timezone=True), nullable=False),
|
||||
sa.Column("created_at", sa.DateTime(timezone=True), server_default=sa.func.now(), nullable=False),
|
||||
sa.ForeignKeyConstraint(["user_id"], ["users.id"], ondelete="CASCADE"),
|
||||
sa.PrimaryKeyConstraint("id"),
|
||||
)
|
||||
op.create_index("ix_password_reset_tokens_user_id", "password_reset_tokens", ["user_id"], unique=False)
|
||||
|
||||
|
||||
def downgrade() -> None:
|
||||
op.drop_index("ix_password_reset_tokens_user_id", table_name="password_reset_tokens")
|
||||
op.drop_table("password_reset_tokens")
|
||||
@@ -12,26 +12,27 @@ dependencies = [
|
||||
# 核心框架
|
||||
"flask>=3.0.0",
|
||||
"flask-cors>=6.0.0",
|
||||
|
||||
# LLM 相关
|
||||
"openai>=1.0.0",
|
||||
|
||||
# Zep Cloud
|
||||
"zep-cloud==3.13.0",
|
||||
|
||||
# OASIS 社交媒体模拟
|
||||
"camel-oasis==0.2.5",
|
||||
"camel-ai==0.2.78",
|
||||
|
||||
# 文件处理
|
||||
"PyMuPDF>=1.24.0",
|
||||
# 编码检测(支持非UTF-8编码的文本文件)
|
||||
"charset-normalizer>=3.0.0",
|
||||
"chardet>=5.0.0",
|
||||
|
||||
# 工具库
|
||||
"python-dotenv>=1.0.0",
|
||||
"pydantic>=2.0.0",
|
||||
"sqlalchemy>=2.0,<3",
|
||||
"gunicorn>=21.2.0",
|
||||
"alembic>=1.13,<2",
|
||||
"argon2-cffi>=23.1",
|
||||
"psycopg[binary]>=3.2",
|
||||
"cryptography>=41.0",
|
||||
]
|
||||
|
||||
[project.optional-dependencies]
|
||||
|
||||
49
backend/tests/fixtures/memory_parity/entity_reader_fixture.json
vendored
Normal file
49
backend/tests/fixtures/memory_parity/entity_reader_fixture.json
vendored
Normal file
@@ -0,0 +1,49 @@
|
||||
{
|
||||
"graph_id": "graph-fixture",
|
||||
"nodes": [
|
||||
{
|
||||
"id": "node-alice",
|
||||
"canonical_name": "Alice",
|
||||
"normalized_name": "alice",
|
||||
"labels": ["Person"],
|
||||
"aliases": [],
|
||||
"attributes": {"role": "founder"},
|
||||
"summary": "A founder."
|
||||
},
|
||||
{
|
||||
"id": "node-acme",
|
||||
"canonical_name": "Acme",
|
||||
"normalized_name": "acme",
|
||||
"labels": ["Organization"],
|
||||
"aliases": [],
|
||||
"attributes": {},
|
||||
"summary": "A company."
|
||||
},
|
||||
{
|
||||
"id": "node-note",
|
||||
"canonical_name": "Launch note",
|
||||
"normalized_name": "launch note",
|
||||
"labels": ["Note"],
|
||||
"aliases": [],
|
||||
"attributes": {},
|
||||
"summary": "A source note."
|
||||
}
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"id": "edge-alice-acme",
|
||||
"source_node_id": "node-alice",
|
||||
"target_node_id": "node-acme",
|
||||
"relation": "WORKS_FOR",
|
||||
"fact": "Alice works for Acme.",
|
||||
"attributes": {},
|
||||
"confidence": 0.9
|
||||
}
|
||||
],
|
||||
"filter": {
|
||||
"defined_entity_types": ["Person", "Organization"],
|
||||
"total_count": 3,
|
||||
"filtered_count": 2,
|
||||
"entity_types": ["Organization", "Person"]
|
||||
}
|
||||
}
|
||||
75
backend/tests/fixtures/memory_parity/tools_fixture.json
vendored
Normal file
75
backend/tests/fixtures/memory_parity/tools_fixture.json
vendored
Normal file
@@ -0,0 +1,75 @@
|
||||
{
|
||||
"graph_id": "graph-tools-fixture",
|
||||
"organization_id": "org-a",
|
||||
"project_id": "project-tools",
|
||||
"nodes": [
|
||||
{
|
||||
"id": "node-alice",
|
||||
"canonical_name": "Alice",
|
||||
"normalized_name": "alice",
|
||||
"labels": ["Entity", "Person"],
|
||||
"aliases": ["A. Example"],
|
||||
"attributes": {"role": "founder"},
|
||||
"summary": "A founder mentioned in the source.",
|
||||
"confidence": 0.91
|
||||
},
|
||||
{
|
||||
"id": "node-acme",
|
||||
"canonical_name": "Acme",
|
||||
"normalized_name": "acme",
|
||||
"labels": ["Entity", "Organization"],
|
||||
"aliases": [],
|
||||
"attributes": {"sector": "software"},
|
||||
"summary": "A software company mentioned in the source.",
|
||||
"confidence": 0.88
|
||||
},
|
||||
{
|
||||
"id": "node-beta",
|
||||
"canonical_name": "Beta",
|
||||
"normalized_name": "beta",
|
||||
"labels": ["Entity", "Organization"],
|
||||
"aliases": [],
|
||||
"attributes": {},
|
||||
"summary": "A partner company mentioned in the source.",
|
||||
"confidence": 0.82
|
||||
}
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"id": "edge-current",
|
||||
"source_node_id": "node-alice",
|
||||
"target_node_id": "node-acme",
|
||||
"relation": "WORKS_FOR",
|
||||
"fact": "Alice works for Acme.",
|
||||
"attributes": {},
|
||||
"confidence": 0.9,
|
||||
"valid_at": "2024-01-01T00:00:00+00:00",
|
||||
"invalid_at": null,
|
||||
"expired_at": null
|
||||
},
|
||||
{
|
||||
"id": "edge-historical",
|
||||
"source_node_id": "node-alice",
|
||||
"target_node_id": "node-acme",
|
||||
"relation": "WORKED_FOR",
|
||||
"fact": "Alice previously worked for Acme.",
|
||||
"attributes": {},
|
||||
"confidence": 0.86,
|
||||
"valid_at": "2020-01-01T00:00:00+00:00",
|
||||
"invalid_at": "2023-01-01T00:00:00+00:00",
|
||||
"expired_at": null
|
||||
},
|
||||
{
|
||||
"id": "edge-unrelated",
|
||||
"source_node_id": "node-acme",
|
||||
"target_node_id": "node-beta",
|
||||
"relation": "PARTNERED_WITH",
|
||||
"fact": "Acme partnered with Beta.",
|
||||
"attributes": {},
|
||||
"confidence": 0.8,
|
||||
"valid_at": "2022-01-01T00:00:00+00:00",
|
||||
"invalid_at": null,
|
||||
"expired_at": null
|
||||
}
|
||||
]
|
||||
}
|
||||
235
backend/tests/test_admin_users_api.py
Normal file
235
backend/tests/test_admin_users_api.py
Normal file
@@ -0,0 +1,235 @@
|
||||
import json
|
||||
|
||||
from flask import Flask
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.api.admin import admin_bp
|
||||
from app.api.auth import auth_bp
|
||||
from app.db import Base, create_session_factory
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
|
||||
|
||||
def make_admin_app():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SECRET_KEY="test-secret", SESSION_COOKIE_SECURE=False)
|
||||
app.extensions["crowdsight_session_factory"] = session_factory
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
app.register_blueprint(admin_bp, url_prefix="/api/admin")
|
||||
|
||||
with session_factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
admin = repo.create_user(
|
||||
email="admin@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(admin.id, org.id, "admin")
|
||||
user = repo.create_user(
|
||||
email="user@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(user.id, org.id, "user")
|
||||
session.commit()
|
||||
|
||||
return app, engine
|
||||
|
||||
|
||||
def login(client, email):
|
||||
response = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": email, "password": "correct horse battery staple"},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
def _csrf_headers(client):
|
||||
return {"X-CSRF-Token": client.get_cookie("crowdsight_csrf").value}
|
||||
|
||||
|
||||
def test_admin_lists_and_creates_users_only_in_current_organization():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
|
||||
listed = client.get("/api/admin/users")
|
||||
assert listed.status_code == 200
|
||||
body = listed.get_json()
|
||||
assert body["data"]["count"] == 2
|
||||
assert all("password_hash" not in user for user in body["data"]["users"])
|
||||
|
||||
created = client.post(
|
||||
"/api/admin/users",
|
||||
json={"email": "new-user@example.com"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert created.status_code == 201
|
||||
assert created.get_json()["data"]["role"] == "user"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_admin_cannot_grant_admin_and_user_cannot_manage_users():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
admin_client = app.test_client()
|
||||
login(admin_client, "admin@example.com")
|
||||
forbidden = admin_client.post(
|
||||
"/api/admin/users",
|
||||
json={"email": "new-admin@example.com", "role": "admin"},
|
||||
headers=_csrf_headers(admin_client),
|
||||
)
|
||||
assert forbidden.status_code == 403
|
||||
assert forbidden.get_json()["error_code"] == "admin_role_grant_forbidden"
|
||||
|
||||
user_client = app.test_client()
|
||||
login(user_client, "user@example.com")
|
||||
user_forbidden = user_client.get("/api/admin/users")
|
||||
assert user_forbidden.status_code == 403
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
def test_admin_duplicate_email_returns_conflict_without_db_error():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
response = client.post(
|
||||
"/api/admin/users",
|
||||
json={"email": "user@example.com"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert response.status_code == 409
|
||||
assert response.get_json()["error_code"] == "user_exists"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
def test_state_changing_admin_request_requires_csrf_token():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
response = client.post(
|
||||
"/api/admin/users",
|
||||
json={"email": "csrf-user@example.com"},
|
||||
)
|
||||
assert response.status_code == 403
|
||||
assert response.get_json()["error_code"] == "csrf_failed"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_admin_can_update_user_role_within_allowed_domain():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
|
||||
# Find the target user's id.
|
||||
listed = client.get("/api/admin/users").get_json()
|
||||
target = next(u for u in listed["data"]["users"] if u["email"] == "user@example.com")
|
||||
|
||||
# Admin may manage a USER role only (policy: user→user is allowed).
|
||||
response = client.patch(
|
||||
f"/api/admin/users/{target['id']}",
|
||||
json={"role": "user"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert response.status_code == 200
|
||||
assert response.get_json()["data"]["role"] == "user"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_admin_cannot_promote_to_super_admin():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
listed = client.get("/api/admin/users").get_json()
|
||||
target = next(u for u in listed["data"]["users"] if u["email"] == "user@example.com")
|
||||
|
||||
response = client.patch(
|
||||
f"/api/admin/users/{target['id']}",
|
||||
json={"role": "super_admin"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert response.status_code == 403
|
||||
assert response.get_json()["error_code"] == "admin_role_grant_forbidden"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_settings_endpoint_requires_super_admin():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
# An org admin (ADMIN) must NOT be able to read or mutate platform settings.
|
||||
forbidden_get = client.get("/api/admin/settings")
|
||||
assert forbidden_get.status_code == 403
|
||||
forbidden_put = client.put(
|
||||
"/api/admin/settings",
|
||||
json={"provider": "openai", "model": "gpt-4o"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert forbidden_put.status_code == 403
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_super_admin_can_save_and_read_settings():
|
||||
app, engine = make_admin_app()
|
||||
try:
|
||||
# SettingsService derives its Fernet key from Config.SECRET_KEY; set it
|
||||
# directly (it was evaluated from env at import time).
|
||||
from app.config import Config
|
||||
|
||||
Config.SECRET_KEY = "test-encryption-secret-" * 3
|
||||
|
||||
# Promote the existing admin membership to super_admin in its org.
|
||||
with app.extensions["crowdsight_session_factory"]() as session:
|
||||
from sqlalchemy import select
|
||||
|
||||
from app.models.saas import Membership, Organization
|
||||
from app.services.identity import IdentityRepository
|
||||
|
||||
repo = IdentityRepository(session)
|
||||
admin_user = repo.get_user_by_email("admin@example.com")
|
||||
org = session.execute(
|
||||
select(Organization).where(Organization.slug == "org-a")
|
||||
).scalar_one()
|
||||
membership = session.execute(
|
||||
select(Membership).where(
|
||||
Membership.user_id == admin_user.id,
|
||||
Membership.organization_id == org.id,
|
||||
)
|
||||
).scalar_one()
|
||||
membership.role = "super_admin"
|
||||
session.commit()
|
||||
|
||||
client = app.test_client()
|
||||
login(client, "admin@example.com")
|
||||
|
||||
saved = client.put(
|
||||
"/api/admin/settings",
|
||||
json={"provider": "openai", "model": "gpt-4o", "api_key": "sk-secret-value"},
|
||||
headers=_csrf_headers(client),
|
||||
)
|
||||
assert saved.status_code == 200
|
||||
data = saved.get_json()["data"]
|
||||
assert data["settings"]["model"] == "gpt-4o"
|
||||
# API key is masked (never plaintext). It may be a bare mask or a
|
||||
# masked-with-prefix variant; the important invariant is no plaintext.
|
||||
assert "sk-secret-value" not in json.dumps(saved.get_json())
|
||||
|
||||
# Reading back returns masked value only; no plaintext secret leaked.
|
||||
got = client.get("/api/admin/settings")
|
||||
assert got.status_code == 200
|
||||
assert "sk-secret-value" not in json.dumps(got.get_json())
|
||||
finally:
|
||||
engine.dispose()
|
||||
48
backend/tests/test_api_errors.py
Normal file
48
backend/tests/test_api_errors.py
Normal file
@@ -0,0 +1,48 @@
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
_MODULE_PATH = Path(__file__).parents[1] / "app" / "utils" / "api_errors.py"
|
||||
_SPEC = importlib.util.spec_from_file_location("api_errors_under_test", _MODULE_PATH)
|
||||
assert _SPEC is not None and _SPEC.loader is not None
|
||||
_MODULE = importlib.util.module_from_spec(_SPEC)
|
||||
sys.modules[_SPEC.name] = _MODULE
|
||||
_SPEC.loader.exec_module(_MODULE)
|
||||
|
||||
ApiError = _MODULE.ApiError
|
||||
internal_error_payload = _MODULE.internal_error_payload
|
||||
|
||||
|
||||
class ApiErrorTests(unittest.TestCase):
|
||||
def translator(self, key, **kwargs):
|
||||
return f"{key}:{kwargs.get('id', '')}".rstrip(":")
|
||||
|
||||
def test_structured_error_has_code_and_localized_message_only(self):
|
||||
error = ApiError(
|
||||
code="project_not_found",
|
||||
status_code=404,
|
||||
message_key="api.projectNotFound",
|
||||
params={"id": "demo"},
|
||||
)
|
||||
payload = error.to_payload(self.translator)
|
||||
self.assertEqual(payload, {
|
||||
"success": False,
|
||||
"error_code": "project_not_found",
|
||||
"message": "api.projectNotFound:demo",
|
||||
})
|
||||
self.assertNotIn("error", payload)
|
||||
self.assertNotIn("details", payload)
|
||||
|
||||
def test_internal_error_never_contains_exception_text(self):
|
||||
payload = internal_error_payload(self.translator)
|
||||
self.assertEqual(payload["success"], False)
|
||||
self.assertEqual(payload["error_code"], "internal_error")
|
||||
self.assertEqual(payload["message"], "api.internalError")
|
||||
self.assertNotIn("Traceback", str(payload))
|
||||
self.assertNotIn("secret", str(payload))
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
36
backend/tests/test_api_no_raw_exception_details.py
Normal file
36
backend/tests/test_api_no_raw_exception_details.py
Normal file
@@ -0,0 +1,36 @@
|
||||
from pathlib import Path
|
||||
import re
|
||||
|
||||
|
||||
API_DIR = Path(__file__).parents[1] / "app" / "api"
|
||||
|
||||
|
||||
UNSAFE_PATTERNS = (
|
||||
re.compile(r"[\"']traceback[\"']\s*:\s*traceback\.format_exc\(\)"),
|
||||
re.compile(r"[\"']error[\"']\s*:\s*str\((?:e|exc)\)"),
|
||||
re.compile(r"t\([^\n]*error\s*=\s*str\((?:e|exc)\)"),
|
||||
re.compile(r"(?:project|state)\.error\s*=\s*str\((?:e|exc)\)"),
|
||||
re.compile(r"fail_task\([^\n]*str\((?:e|exc)\)"),
|
||||
)
|
||||
|
||||
|
||||
TARGET_FILES = (
|
||||
"simulation.py",
|
||||
"graph.py",
|
||||
"report.py",
|
||||
"agent_group.py",
|
||||
"template.py",
|
||||
)
|
||||
|
||||
|
||||
def test_api_does_not_expose_or_persist_raw_exception_details():
|
||||
findings = []
|
||||
for filename in TARGET_FILES:
|
||||
path = API_DIR / filename
|
||||
source = path.read_text(encoding="utf-8")
|
||||
for pattern in UNSAFE_PATTERNS:
|
||||
for match in pattern.finditer(source):
|
||||
line = source.count("\n", 0, match.start()) + 1
|
||||
findings.append(f"{filename}:{line}:{match.group(0)}")
|
||||
|
||||
assert findings == [], "unsafe exception details found:\n" + "\n".join(findings)
|
||||
70
backend/tests/test_artifact_store.py
Normal file
70
backend/tests/test_artifact_store.py
Normal file
@@ -0,0 +1,70 @@
|
||||
"""TDD gate: tenant-scoped artifact store abstraction.
|
||||
|
||||
The legacy project/simulation/report managers write to raw os.path.join under a
|
||||
shared upload root. This gate proves a scoped ArtifactStore resolves paths
|
||||
within the tenant's own directory, rejects traversal/absolute components, and
|
||||
exposes a store/read/delete interface that can later be backed by object
|
||||
storage without changing callers.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
|
||||
from app.services.artifact_store import ArtifactStore
|
||||
|
||||
|
||||
def test_artifact_store_scopes_path_to_root(monkeypatch, tmp_path):
|
||||
store = ArtifactStore(str(tmp_path))
|
||||
project_dir = store.path_for("org-a", "project-1")
|
||||
assert str(tmp_path) in project_dir
|
||||
assert "org-a" in project_dir
|
||||
assert "project-1" in project_dir
|
||||
# Tenant-scoped directory lands under the configured root.
|
||||
assert project_dir.startswith(str(tmp_path))
|
||||
|
||||
|
||||
def test_artifact_store_isolates_tenants(monkeypatch, tmp_path):
|
||||
store = ArtifactStore(str(tmp_path))
|
||||
path_a = store.path_for("org-a", "proj-x", "state.json")
|
||||
path_b = store.path_for("org-b", "proj-x", "state.json")
|
||||
assert path_a != path_b
|
||||
assert "/org-a/" in path_a.replace("\\", "/")
|
||||
assert "/org-b/" in path_b.replace("\\", "/")
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
"bad_segment",
|
||||
["../evil", "a/../../etc", "/absolute", "..", ".../..", "a//..", "\\evil", "a/..\\b"],
|
||||
)
|
||||
def test_artifact_store_rejects_traversal(bad_segment, tmp_path):
|
||||
store = ArtifactStore(str(tmp_path))
|
||||
with pytest.raises(ValueError):
|
||||
store.path_for("org-a", bad_segment)
|
||||
|
||||
|
||||
def test_artifact_store_roundtrip_bytes(monkeypatch, tmp_path):
|
||||
store = ArtifactStore(str(tmp_path))
|
||||
target = store.path_for("org-a", "report-1", "report.md")
|
||||
store.ensure_parent(target)
|
||||
store.store_bytes(target, b"# title\nbody")
|
||||
assert store.read_bytes(target) == b"# title\nbody"
|
||||
assert store.exists(target)
|
||||
store.delete(target)
|
||||
assert not store.exists(target)
|
||||
|
||||
|
||||
def test_default_artifact_store_uses_configured_upload_root(monkeypatch, tmp_path):
|
||||
from app.config import Config
|
||||
from app.services.artifact_store import default_artifact_store
|
||||
|
||||
original = Config.UPLOAD_FOLDER
|
||||
monkeypatch.setattr(Config, "UPLOAD_FOLDER", str(tmp_path / "uploads"))
|
||||
try:
|
||||
store = default_artifact_store()
|
||||
project_dir = store.path_for("org-a", "project-1")
|
||||
assert str(Config.UPLOAD_FOLDER) in project_dir
|
||||
# Tenant-scoped and confinement still apply.
|
||||
assert ".." not in project_dir
|
||||
with pytest.raises(ValueError):
|
||||
store.path_for("org-a", "..")
|
||||
finally:
|
||||
Config.UPLOAD_FOLDER = original
|
||||
127
backend/tests/test_audit_service.py
Normal file
127
backend/tests/test_audit_service.py
Normal file
@@ -0,0 +1,127 @@
|
||||
"""TDD gate: durable, redacted audit event recording.
|
||||
|
||||
Audit events must be tenant-scoped and never store secrets, tokens, password
|
||||
hashes, or raw prompts. This gate proves an AuditService that records actions
|
||||
and lists them scoped to an organization.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.services.audit_service import AuditService
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def session_factory():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
try:
|
||||
yield factory
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_audit_service_records_and_lists_scoped(session_factory):
|
||||
session = session_factory()
|
||||
try:
|
||||
svc = AuditService(session)
|
||||
entry_id = svc.record(
|
||||
organization_id="org-a",
|
||||
actor_user_id="user-1",
|
||||
action="user.role_changed",
|
||||
target_type="user",
|
||||
target_id="user-2",
|
||||
details={"role": "admin"},
|
||||
)
|
||||
assert entry_id
|
||||
|
||||
rows = svc.list_for_organization(organization_id="org-a")
|
||||
assert len(rows) == 1
|
||||
assert rows[0].action == "user.role_changed"
|
||||
assert rows[0].actor_user_id == "user-1"
|
||||
|
||||
# Scoped list from another org does not see it.
|
||||
other = svc.list_for_organization(organization_id="org-b")
|
||||
assert other == []
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_audit_service_redacts_secrets_from_details(session_factory):
|
||||
session = session_factory()
|
||||
try:
|
||||
svc = AuditService(session)
|
||||
svc.record(
|
||||
organization_id="org-a",
|
||||
actor_user_id="user-1",
|
||||
action="auth.password_changed",
|
||||
target_type="user",
|
||||
target_id="user-2",
|
||||
details={"password": "hunter2", "token": "abc", "api_key": "sk-xyz"},
|
||||
)
|
||||
row = svc.list_for_organization(organization_id="org-a")[0]
|
||||
details = row.details if isinstance(row.details, dict) else {}
|
||||
blob = repr(details)
|
||||
assert "hunter2" not in blob
|
||||
assert "abc" not in blob
|
||||
assert "sk-xyz" not in blob
|
||||
# Redaction leaves non-sensitive context intact.
|
||||
assert details == {}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_login_endpoint_writes_audit_event(tmp_path, monkeypatch, session_factory):
|
||||
"""A successful login records an auth.login audit event."""
|
||||
import secrets
|
||||
|
||||
from flask import Flask, jsonify
|
||||
|
||||
from app.api.auth import auth_bp
|
||||
from app.services.audit_service import AuditService
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
from app.utils.api_errors import ApiError
|
||||
from app.utils.locale import t
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SESSION_COOKIE_SECURE=False)
|
||||
app.config["SECRET_KEY"] = secrets.token_hex(32)
|
||||
app.extensions["crowdsight_session_factory"] = factory
|
||||
|
||||
@app.errorhandler(ApiError)
|
||||
def handle_api_error(error):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
|
||||
audit_email = "audit-login@example.com"
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Audit Org", slug="audit-org")
|
||||
user = repo.create_user(
|
||||
email=audit_email,
|
||||
password_hash=PasswordService.hash_password("correct-horse"),
|
||||
)
|
||||
repo.create_membership(user.id, org.id, "user")
|
||||
session.commit()
|
||||
user_id = user.id
|
||||
org_id = org.id
|
||||
|
||||
client = app.test_client()
|
||||
r = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": audit_email, "password": "correct-horse"},
|
||||
)
|
||||
assert r.status_code == 200
|
||||
|
||||
# The audit event exists and is redacted/org-scoped.
|
||||
with factory() as session:
|
||||
events = AuditService(session).list_for_organization(organization_id=org_id)
|
||||
assert any(e.action == "auth.login" and e.actor_user_id == user_id for e in events)
|
||||
engine.dispose()
|
||||
81
backend/tests/test_auth_api.py
Normal file
81
backend/tests/test_auth_api.py
Normal file
@@ -0,0 +1,81 @@
|
||||
import json
|
||||
|
||||
from flask import Flask
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
|
||||
from app.api.auth import auth_bp
|
||||
from app.db import Base, create_session_factory
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
|
||||
|
||||
def make_auth_app():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SECRET_KEY="test-secret", SESSION_COOKIE_SECURE=False)
|
||||
app.extensions["crowdsight_session_factory"] = session_factory
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
|
||||
with session_factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
user = repo.create_user(
|
||||
email="admin@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(user.id, org.id, "admin")
|
||||
session.commit()
|
||||
|
||||
return app, engine
|
||||
|
||||
|
||||
def _csrf_headers(client):
|
||||
return {"X-CSRF-Token": client.get_cookie("crowdsight_csrf").value}
|
||||
|
||||
|
||||
def test_login_me_logout_uses_cookie_and_allowlisted_identity():
|
||||
app, engine = make_auth_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
login = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": "ADMIN@example.com", "password": "correct horse battery staple"},
|
||||
)
|
||||
assert login.status_code == 200
|
||||
body = login.get_json()
|
||||
assert body["success"] is True
|
||||
assert body["data"]["user"]["email"] == "admin@example.com"
|
||||
assert body["data"]["role"] == "admin"
|
||||
assert body["data"]["organization"]["slug"] == "org-a"
|
||||
assert "token" not in json.dumps(body)
|
||||
assert "password" not in json.dumps(body).lower()
|
||||
|
||||
me = client.get("/api/auth/me")
|
||||
assert me.status_code == 200
|
||||
assert me.get_json()["data"]["user"]["id"]
|
||||
|
||||
logout = client.post("/api/auth/logout", headers=_csrf_headers(client))
|
||||
assert logout.status_code == 200
|
||||
assert client.get("/api/auth/me").status_code == 401
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_invalid_credentials_return_structured_generic_error():
|
||||
app, engine = make_auth_app()
|
||||
try:
|
||||
response = app.test_client().post(
|
||||
"/api/auth/login",
|
||||
json={"email": "admin@example.com", "password": "wrong password"},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
body = response.get_json()
|
||||
assert body["success"] is False
|
||||
assert body["error_code"] == "invalid_credentials"
|
||||
assert "error" not in body
|
||||
assert "Traceback" not in json.dumps(body)
|
||||
finally:
|
||||
engine.dispose()
|
||||
87
backend/tests/test_auth_service.py
Normal file
87
backend/tests/test_auth_service.py
Normal file
@@ -0,0 +1,87 @@
|
||||
from datetime import datetime, timedelta, timezone
|
||||
|
||||
from app.models.saas import AuthSession
|
||||
from app.services.identity import IdentityRepository, PasswordService, SessionService
|
||||
|
||||
|
||||
def test_session_token_is_hashed_and_resolvable():
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
from app.db import Base
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
try:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
user = repo.create_user(
|
||||
email="admin@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
membership = repo.create_membership(user.id, org.id, "admin")
|
||||
raw_token, stored = SessionService.create(session, user, membership.id)
|
||||
session.commit()
|
||||
|
||||
assert raw_token != stored.token_hash
|
||||
assert len(stored.token_hash) == 64
|
||||
resolved = SessionService.resolve(session, raw_token)
|
||||
assert resolved is not None
|
||||
assert resolved.user.id == user.id
|
||||
assert resolved.membership.id == membership.id
|
||||
assert resolved.organization.id == org.id
|
||||
finally:
|
||||
session.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_revoked_or_stale_auth_version_session_is_not_resolvable():
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
from app.db import Base
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
try:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
user = repo.create_user(email="user@example.com")
|
||||
membership = repo.create_membership(user.id, org.id, "user")
|
||||
raw_token, _stored = SessionService.create(session, user, membership.id)
|
||||
session.commit()
|
||||
|
||||
user.auth_version += 1
|
||||
session.commit()
|
||||
assert SessionService.resolve(session, raw_token) is None
|
||||
|
||||
user.auth_version -= 1
|
||||
session.commit()
|
||||
assert SessionService.revoke(session, raw_token) is True
|
||||
session.commit()
|
||||
assert SessionService.resolve(session, raw_token) is None
|
||||
finally:
|
||||
session.close()
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_expired_session_is_not_resolvable():
|
||||
from sqlalchemy import create_engine
|
||||
from sqlalchemy.orm import Session
|
||||
from app.db import Base
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session = Session(engine)
|
||||
try:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
user = repo.create_user(email="expired@example.com")
|
||||
membership = repo.create_membership(user.id, org.id, "user")
|
||||
raw_token, stored = SessionService.create(session, user, membership.id)
|
||||
stored.expires_at = datetime.now(timezone.utc) - timedelta(seconds=1)
|
||||
session.commit()
|
||||
assert SessionService.resolve(session, raw_token) is None
|
||||
finally:
|
||||
session.close()
|
||||
engine.dispose()
|
||||
126
backend/tests/test_authorization_policy.py
Normal file
126
backend/tests/test_authorization_policy.py
Normal file
@@ -0,0 +1,126 @@
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
_MODULE_PATH = Path(__file__).parents[1] / "app" / "security" / "policy.py"
|
||||
_SPEC = importlib.util.spec_from_file_location("authorization_policy_under_test", _MODULE_PATH)
|
||||
assert _SPEC is not None and _SPEC.loader is not None
|
||||
_POLICY = importlib.util.module_from_spec(_SPEC)
|
||||
sys.modules[_SPEC.name] = _POLICY
|
||||
_SPEC.loader.exec_module(_POLICY)
|
||||
|
||||
Actor = _POLICY.Actor
|
||||
AuthorizationError = _POLICY.AuthorizationError
|
||||
Role = _POLICY.Role
|
||||
assert_can_access_resource = _POLICY.assert_can_access_resource
|
||||
assert_can_manage_llm_settings = _POLICY.assert_can_manage_llm_settings
|
||||
assert_can_manage_user = _POLICY.assert_can_manage_user
|
||||
|
||||
|
||||
class AuthorizationPolicyTests(unittest.TestCase):
|
||||
def setUp(self):
|
||||
self.admin = Actor(user_id="admin-1", organization_id="org-a", role=Role.ADMIN)
|
||||
self.user = Actor(user_id="user-1", organization_id="org-a", role=Role.USER)
|
||||
self.other_user = Actor(user_id="user-2", organization_id="org-a", role=Role.USER)
|
||||
self.super_admin = Actor(
|
||||
user_id="root-1", organization_id="platform", role=Role.SUPER_ADMIN
|
||||
)
|
||||
|
||||
def test_user_can_access_owned_resource_in_own_organization(self):
|
||||
assert_can_access_resource(
|
||||
self.user,
|
||||
resource_organization_id="org-a",
|
||||
owner_user_id="user-1",
|
||||
)
|
||||
|
||||
def test_user_cannot_access_another_users_resource_in_same_organization(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_access_resource(
|
||||
self.user,
|
||||
resource_organization_id="org-a",
|
||||
owner_user_id="user-2",
|
||||
)
|
||||
|
||||
def test_user_cannot_access_resource_from_another_organization(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_access_resource(
|
||||
self.user,
|
||||
resource_organization_id="org-b",
|
||||
owner_user_id="user-1",
|
||||
)
|
||||
|
||||
def test_admin_can_access_resources_in_own_organization(self):
|
||||
assert_can_access_resource(
|
||||
self.admin,
|
||||
resource_organization_id="org-a",
|
||||
owner_user_id="user-2",
|
||||
)
|
||||
|
||||
def test_admin_cannot_cross_organization_boundary(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_access_resource(
|
||||
self.admin,
|
||||
resource_organization_id="org-b",
|
||||
owner_user_id="user-9",
|
||||
)
|
||||
|
||||
def test_super_admin_requires_explicit_platform_scope_for_cross_tenant_access(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_access_resource(
|
||||
self.super_admin,
|
||||
resource_organization_id="org-a",
|
||||
owner_user_id="user-9",
|
||||
)
|
||||
|
||||
assert_can_access_resource(
|
||||
self.super_admin,
|
||||
resource_organization_id="org-a",
|
||||
owner_user_id="user-9",
|
||||
platform_scope=True,
|
||||
)
|
||||
|
||||
def test_admin_can_manage_user_but_not_grant_admin_or_super_admin(self):
|
||||
assert_can_manage_user(
|
||||
self.admin,
|
||||
target_organization_id="org-a",
|
||||
target_role=Role.USER,
|
||||
)
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_manage_user(
|
||||
self.admin,
|
||||
target_organization_id="org-a",
|
||||
target_role=Role.ADMIN,
|
||||
)
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_manage_user(
|
||||
self.admin,
|
||||
target_organization_id="org-a",
|
||||
target_role=Role.SUPER_ADMIN,
|
||||
)
|
||||
|
||||
def test_super_admin_can_manage_any_role_only_with_platform_scope(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_manage_user(
|
||||
self.super_admin,
|
||||
target_organization_id="org-a",
|
||||
target_role=Role.ADMIN,
|
||||
)
|
||||
assert_can_manage_user(
|
||||
self.super_admin,
|
||||
target_organization_id="org-a",
|
||||
target_role=Role.ADMIN,
|
||||
platform_scope=True,
|
||||
)
|
||||
|
||||
def test_only_super_admin_can_manage_llm_settings(self):
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_manage_llm_settings(self.admin)
|
||||
with self.assertRaises(AuthorizationError):
|
||||
assert_can_manage_llm_settings(self.user)
|
||||
assert_can_manage_llm_settings(self.super_admin, platform_scope=True)
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
249
backend/tests/test_auxiliary_api_security.py
Normal file
249
backend/tests/test_auxiliary_api_security.py
Normal file
@@ -0,0 +1,249 @@
|
||||
from __future__ import annotations
|
||||
|
||||
import importlib
|
||||
import secrets
|
||||
|
||||
import pytest
|
||||
from flask import Flask, jsonify
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.api.agent_group import agent_group_bp
|
||||
from app.api.auth import auth_bp
|
||||
from app.api.template import template_bp
|
||||
from app.db import Base, create_session_factory
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
from app.utils.api_errors import ApiError
|
||||
from app.utils.locale import t
|
||||
|
||||
agent_group_module = importlib.import_module("app.api.agent_group")
|
||||
template_module = importlib.import_module("app.api.template")
|
||||
|
||||
|
||||
TEST_EMAIL = "auxiliary-security@example.com"
|
||||
TEST_AUTH_INPUT = "local-only-auth-input"
|
||||
|
||||
|
||||
def make_app():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
|
||||
app = Flask(__name__)
|
||||
app.config.update(
|
||||
TESTING=True,
|
||||
SESSION_COOKIE_SECURE=False,
|
||||
)
|
||||
app.config["SECRET_KEY"] = secrets.token_hex(32)
|
||||
app.extensions["crowdsight_session_factory"] = session_factory
|
||||
|
||||
@app.errorhandler(ApiError)
|
||||
def handle_api_error(error: ApiError):
|
||||
return jsonify(error.to_payload(t)), error.status_code
|
||||
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
app.register_blueprint(template_bp, url_prefix="/api/template")
|
||||
app.register_blueprint(agent_group_bp, url_prefix="/api/agent-group")
|
||||
|
||||
with session_factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
organization = repo.create_organization(name="Auxiliary Org", slug="auxiliary-org")
|
||||
user = repo.create_user(
|
||||
email=TEST_EMAIL,
|
||||
password_hash=PasswordService.hash_password(TEST_AUTH_INPUT),
|
||||
)
|
||||
repo.create_membership(user.id, organization.id, "admin")
|
||||
session.commit()
|
||||
|
||||
return app, engine
|
||||
|
||||
|
||||
def login(client):
|
||||
response = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": TEST_EMAIL, "password": TEST_AUTH_INPUT},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
return client.get_cookie("crowdsight_csrf").value
|
||||
|
||||
|
||||
def test_template_list_requires_authentication():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
response = app.test_client().get("/api/template/list")
|
||||
assert response.status_code == 401
|
||||
assert response.get_json()["error_code"] == "unauthorized"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_agent_group_categorize_requires_authentication():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
response = app.test_client().post(
|
||||
"/api/agent-group/categorize",
|
||||
json={"agents": [{"name": "Alice"}]},
|
||||
)
|
||||
assert response.status_code == 401
|
||||
assert response.get_json()["error_code"] == "unauthorized"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_authenticated_auxiliary_reads_and_pure_filter_remain_available():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
csrf = login(client)
|
||||
|
||||
templates = client.get("/api/template/list")
|
||||
assert templates.status_code == 200
|
||||
assert templates.get_json()["success"] is True
|
||||
|
||||
filtered = client.post(
|
||||
"/api/agent-group/filter",
|
||||
json={
|
||||
"agents": [{"agent_id": 1}],
|
||||
"groups": [{"group_id": "all", "agent_indices": [0]}],
|
||||
"selected_group_ids": ["all"],
|
||||
},
|
||||
headers={"X-CSRF-Token": csrf},
|
||||
)
|
||||
assert filtered.status_code == 200
|
||||
assert filtered.get_json()["selected_agent_ids"] == [0]
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
@pytest.mark.parametrize(
|
||||
("path", "payload"),
|
||||
[
|
||||
("/api/template/auto-select", {"text": "Alice founded Orbit."}),
|
||||
("/api/agent-group/categorize", {"agents": [{"name": "Alice"}]}),
|
||||
],
|
||||
)
|
||||
def test_auxiliary_llm_mutations_require_csrf_and_idempotency(path, payload):
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
csrf = login(client)
|
||||
|
||||
missing_csrf = client.post(
|
||||
path,
|
||||
json=payload,
|
||||
headers={"Idempotency-Key": "auxiliary-mutation-1"},
|
||||
)
|
||||
assert missing_csrf.status_code == 403
|
||||
assert missing_csrf.get_json()["error_code"] == "csrf_failed"
|
||||
|
||||
missing_idempotency = client.post(
|
||||
path,
|
||||
json=payload,
|
||||
headers={"X-CSRF-Token": csrf},
|
||||
)
|
||||
assert missing_idempotency.status_code == 400
|
||||
assert missing_idempotency.get_json()["error_code"] == "idempotency_required"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_every_auxiliary_route_requires_authentication():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
routes = [
|
||||
("GET", "/api/template/list", None),
|
||||
("GET", "/api/template/news_event/filter-rules", None),
|
||||
("POST", "/api/template/auto-select", {"text": "Alice founded Orbit."}),
|
||||
("POST", "/api/agent-group/filter", {"agents": [], "groups": []}),
|
||||
("POST", "/api/agent-group/categorize", {"agents": [{"name": "Alice"}]}),
|
||||
]
|
||||
|
||||
for method, path, payload in routes:
|
||||
response = client.open(path, method=method, json=payload)
|
||||
assert response.status_code == 401, (method, path, response.get_json())
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_llm_auxiliary_errors_are_safe(monkeypatch):
|
||||
class ExplodingLLM:
|
||||
def chat_json(self, **_kwargs):
|
||||
raise RuntimeError("sensitive backend detail")
|
||||
|
||||
monkeypatch.setattr(template_module, "LLMClient", ExplodingLLM)
|
||||
monkeypatch.setattr(agent_group_module, "LLMClient", ExplodingLLM)
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
csrf = login(client)
|
||||
responses = [
|
||||
client.post(
|
||||
"/api/template/auto-select",
|
||||
json={"text": "Alice founded Orbit."},
|
||||
headers={"X-CSRF-Token": csrf, "Idempotency-Key": "safe-template-error"},
|
||||
),
|
||||
client.post(
|
||||
"/api/agent-group/categorize",
|
||||
json={"agents": [{"name": "Alice"}]},
|
||||
headers={"X-CSRF-Token": csrf, "Idempotency-Key": "safe-agent-error"},
|
||||
),
|
||||
]
|
||||
for response in responses:
|
||||
body = response.get_json()
|
||||
assert response.status_code == 500
|
||||
assert body["error_code"] == "internal_error"
|
||||
assert "sensitive backend detail" not in response.get_data(as_text=True)
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_llm_auxiliary_mutations_replay_completed_responses(monkeypatch):
|
||||
class FakeTemplateLLM:
|
||||
def chat_json(self, **_kwargs):
|
||||
return {
|
||||
"template_id": "news_event",
|
||||
"prompt": "Alice founded Orbit.",
|
||||
"confidence": 0.9,
|
||||
"reasoning": "The input describes a news event.",
|
||||
}
|
||||
|
||||
class FakeAgentLLM:
|
||||
def chat_json(self, **_kwargs):
|
||||
return {
|
||||
"groups": [
|
||||
{
|
||||
"group_id": "audience",
|
||||
"group_name": "Audience",
|
||||
"default_enabled": True,
|
||||
"agent_indices": [0],
|
||||
}
|
||||
]
|
||||
}
|
||||
|
||||
monkeypatch.setattr(template_module, "LLMClient", FakeTemplateLLM)
|
||||
monkeypatch.setattr(agent_group_module, "LLMClient", FakeAgentLLM)
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
csrf = login(client)
|
||||
cases = [
|
||||
(
|
||||
"/api/template/auto-select",
|
||||
{"text": "Alice founded Orbit."},
|
||||
"template-replay",
|
||||
),
|
||||
(
|
||||
"/api/agent-group/categorize",
|
||||
{"agents": [{"name": "Alice"}]},
|
||||
"agent-replay",
|
||||
),
|
||||
]
|
||||
for path, payload, key in cases:
|
||||
headers = {"X-CSRF-Token": csrf, "Idempotency-Key": key}
|
||||
first = client.post(path, json=payload, headers=headers)
|
||||
second = client.post(path, json=payload, headers=headers)
|
||||
assert first.status_code == 200
|
||||
assert second.status_code == 200
|
||||
assert second.get_json() == first.get_json()
|
||||
finally:
|
||||
engine.dispose()
|
||||
41
backend/tests/test_child_resource_scope.py
Normal file
41
backend/tests/test_child_resource_scope.py
Normal file
@@ -0,0 +1,41 @@
|
||||
from tempfile import TemporaryDirectory
|
||||
|
||||
from app.models.project import ProjectManager
|
||||
from app.services.simulation_manager import SimulationManager
|
||||
from test_resource_auth_scope import csrf_headers, login, make_resource_app
|
||||
|
||||
|
||||
def test_simulation_create_rejects_in_scope_graph_from_another_project():
|
||||
app, engine, organization_id, user_a_id, _user_b_id = make_resource_app()
|
||||
original_projects_dir = ProjectManager.PROJECTS_DIR
|
||||
original_simulations_dir = SimulationManager.SIMULATION_DATA_DIR
|
||||
try:
|
||||
with TemporaryDirectory() as temp_dir:
|
||||
ProjectManager.PROJECTS_DIR = f"{temp_dir}/projects"
|
||||
SimulationManager.SIMULATION_DATA_DIR = f"{temp_dir}/simulations"
|
||||
project_a = ProjectManager.create_project(
|
||||
"A", organization_id=organization_id, owner_user_id=user_a_id
|
||||
)
|
||||
project_a.graph_id = "graph-a"
|
||||
ProjectManager.save_project(project_a)
|
||||
project_b = ProjectManager.create_project(
|
||||
"B", organization_id=organization_id, owner_user_id=user_a_id
|
||||
)
|
||||
project_b.graph_id = "graph-b"
|
||||
ProjectManager.save_project(project_b)
|
||||
|
||||
client = app.test_client()
|
||||
login(client)
|
||||
response = client.post(
|
||||
"/api/simulation/create",
|
||||
json={"project_id": project_a.project_id, "graph_id": project_b.graph_id},
|
||||
headers={
|
||||
**csrf_headers(client),
|
||||
"Idempotency-Key": "cross-project-graph",
|
||||
},
|
||||
)
|
||||
assert response.status_code == 404
|
||||
finally:
|
||||
ProjectManager.PROJECTS_DIR = original_projects_dir
|
||||
SimulationManager.SIMULATION_DATA_DIR = original_simulations_dir
|
||||
engine.dispose()
|
||||
43
backend/tests/test_cors_config.py
Normal file
43
backend/tests/test_cors_config.py
Normal file
@@ -0,0 +1,43 @@
|
||||
import os
|
||||
|
||||
import pytest
|
||||
|
||||
from app import create_app
|
||||
from app.config import Config
|
||||
from app.models.task import TaskManager
|
||||
|
||||
|
||||
class TestConfig(Config):
|
||||
SECRET_KEY = "test-secret"
|
||||
MEMORY_BACKEND = "local"
|
||||
CORS_ALLOWED_ORIGINS = ["https://allowed.example"]
|
||||
TESTING = True
|
||||
|
||||
|
||||
def test_cors_uses_app_allowlist_and_credentials():
|
||||
previous = os.environ.get("DATABASE_URL")
|
||||
os.environ["DATABASE_URL"] = "sqlite+pysqlite:///:memory:"
|
||||
try:
|
||||
app = create_app(TestConfig)
|
||||
response = app.test_client().options(
|
||||
"/api/graph/project/list",
|
||||
headers={
|
||||
"Origin": "https://allowed.example",
|
||||
"Access-Control-Request-Method": "GET",
|
||||
},
|
||||
)
|
||||
assert response.headers["Access-Control-Allow-Origin"] == "https://allowed.example"
|
||||
assert response.headers["Access-Control-Allow-Credentials"] == "true"
|
||||
finally:
|
||||
if previous is None:
|
||||
os.environ.pop("DATABASE_URL", None)
|
||||
else:
|
||||
os.environ["DATABASE_URL"] = previous
|
||||
|
||||
|
||||
def test_cors_wildcard_is_rejected_with_cookie_auth():
|
||||
class WildcardConfig(TestConfig):
|
||||
CORS_ALLOWED_ORIGINS = ["*"]
|
||||
|
||||
with pytest.raises(RuntimeError, match="wildcard_cors_not_allowed"):
|
||||
create_app(WildcardConfig)
|
||||
209
backend/tests/test_durable_task_manager.py
Normal file
209
backend/tests/test_durable_task_manager.py
Normal file
@@ -0,0 +1,209 @@
|
||||
import pytest
|
||||
from flask import Flask
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.task import TaskManager, TaskStatus
|
||||
from app.security.policy import Role
|
||||
from app.services.identity import IdentityRepository
|
||||
|
||||
|
||||
def test_task_state_survives_task_manager_reconfiguration():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
try:
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Jobs", slug="jobs-org")
|
||||
user = repo.create_user(email="jobs@example.com")
|
||||
repo.create_membership(user.id, org.id, Role.ADMIN)
|
||||
session.commit()
|
||||
|
||||
TaskManager.configure(factory)
|
||||
first_manager = TaskManager()
|
||||
task_id = first_manager.create_task(
|
||||
"durable_test",
|
||||
metadata={"organization_id": org.id, "owner_user_id": user.id},
|
||||
)
|
||||
first_manager.update_task(
|
||||
task_id,
|
||||
status=TaskStatus.PROCESSING,
|
||||
progress=42,
|
||||
message="working",
|
||||
progress_detail={"stage": "one"},
|
||||
)
|
||||
|
||||
TaskManager.configure(factory)
|
||||
restarted_manager = TaskManager()
|
||||
task = restarted_manager.get_task(task_id)
|
||||
assert task is not None
|
||||
assert task.status is TaskStatus.PROCESSING
|
||||
assert task.progress == 42
|
||||
assert task.message == "working"
|
||||
assert task.progress_detail == {"stage": "one"}
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_task_manager_query_filters_are_tenant_scoped():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
try:
|
||||
TaskManager.configure(factory)
|
||||
manager = TaskManager()
|
||||
task_a = manager.create_task("alpha", {"organization_id": "org-a", "owner_user_id": "user-a"})
|
||||
task_b = manager.create_task("beta", {"organization_id": "org-b", "owner_user_id": "user-b"})
|
||||
|
||||
assert manager.get_task(task_a, organization_id="org-a") is not None
|
||||
assert manager.get_task(task_b, organization_id="org-a") is None
|
||||
assert manager.get_task(task_a, organization_id="org-a", owner_user_id="user-b") is None
|
||||
assert [task.metadata["organization_id"] for task in manager.list_tasks(organization_id="org-a")] == ["org-a"]
|
||||
assert [task.metadata["owner_user_id"] for task in manager.list_tasks(organization_id="org-a", owner_user_id="user-a")] == ["user-a"]
|
||||
assert manager.list_tasks(organization_id="org-a", owner_user_id="user-b") == []
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_task_manager_fails_closed_inside_app_without_session_factory():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
app = Flask("missing-session-factory")
|
||||
try:
|
||||
TaskManager.configure(factory)
|
||||
with app.app_context():
|
||||
with pytest.raises(RuntimeError, match="task_session_factory_required"):
|
||||
TaskManager()
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_task_manager_binds_to_current_app_without_cross_app_leak():
|
||||
engines = []
|
||||
apps = []
|
||||
organization_ids = []
|
||||
try:
|
||||
for slug in ("app-a", "app-b"):
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
organization = repo.create_organization(name=slug, slug=slug)
|
||||
session.commit()
|
||||
app = Flask(slug)
|
||||
app.extensions["crowdsight_session_factory"] = factory
|
||||
engines.append(engine)
|
||||
apps.append(app)
|
||||
organization_ids.append(organization.id)
|
||||
|
||||
TaskManager.configure(None)
|
||||
with apps[0].app_context():
|
||||
task_a = TaskManager().create_task(
|
||||
"app_a_task",
|
||||
metadata={"organization_id": organization_ids[0]},
|
||||
)
|
||||
with apps[1].app_context():
|
||||
assert TaskManager().get_task(task_a) is None
|
||||
task_b = TaskManager().create_task(
|
||||
"app_b_task",
|
||||
metadata={"organization_id": organization_ids[1]},
|
||||
)
|
||||
with apps[0].app_context():
|
||||
assert TaskManager().get_task(task_a) is not None
|
||||
assert TaskManager().get_task(task_b) is None
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
for engine in engines:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_reused_task_manager_fails_closed_across_app_contexts():
|
||||
engines = []
|
||||
apps = []
|
||||
organization_ids = []
|
||||
try:
|
||||
for slug in ("reuse-a", "reuse-b"):
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
with factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
organization = repo.create_organization(name=slug, slug=slug)
|
||||
session.commit()
|
||||
app = Flask(slug)
|
||||
app.extensions["crowdsight_session_factory"] = factory
|
||||
engines.append(engine)
|
||||
apps.append(app)
|
||||
organization_ids.append(organization.id)
|
||||
|
||||
with apps[0].app_context():
|
||||
manager = TaskManager()
|
||||
task_id = manager.create_task("reuse_task", {"organization_id": organization_ids[0]})
|
||||
with apps[1].app_context():
|
||||
with pytest.raises(RuntimeError, match="task_app_context_mismatch"):
|
||||
manager.get_task(task_id)
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
for engine in engines:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_explicit_task_manager_factory_must_match_current_app():
|
||||
engine_a = create_engine("sqlite+pysqlite:///:memory:")
|
||||
engine_b = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine_a)
|
||||
Base.metadata.create_all(engine_b)
|
||||
factory_a = create_session_factory(engine_a)
|
||||
factory_b = create_session_factory(engine_b)
|
||||
app_b = Flask("explicit-b")
|
||||
app_b.extensions["crowdsight_session_factory"] = factory_b
|
||||
try:
|
||||
with app_b.app_context():
|
||||
with pytest.raises(RuntimeError, match="task_session_factory_mismatch"):
|
||||
TaskManager(session_factory=factory_a)
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine_a.dispose()
|
||||
engine_b.dispose()
|
||||
|
||||
|
||||
def test_prebound_task_manager_fails_closed_inside_other_app():
|
||||
engine_a = create_engine("sqlite+pysqlite:///:memory:")
|
||||
engine_b = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine_a)
|
||||
Base.metadata.create_all(engine_b)
|
||||
factory_a = create_session_factory(engine_a)
|
||||
factory_b = create_session_factory(engine_b)
|
||||
app_b = Flask("prebound-b")
|
||||
app_b.extensions["crowdsight_session_factory"] = factory_b
|
||||
manager = TaskManager(session_factory=factory_a)
|
||||
try:
|
||||
with app_b.app_context():
|
||||
with pytest.raises(RuntimeError, match="task_session_factory_mismatch"):
|
||||
manager.get_task("not-in-scope")
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine_a.dispose()
|
||||
engine_b.dispose()
|
||||
|
||||
|
||||
def test_app_bound_task_manager_supports_background_use_without_context():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
app = Flask("background-use")
|
||||
app.extensions["crowdsight_session_factory"] = factory
|
||||
try:
|
||||
with app.app_context():
|
||||
manager = TaskManager()
|
||||
task_id = manager.create_task("background_task", {"organization_id": "org-bg"})
|
||||
assert manager.get_task(task_id) is not None
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine.dispose()
|
||||
101
backend/tests/test_graph_auth_scope.py
Normal file
101
backend/tests/test_graph_auth_scope.py
Normal file
@@ -0,0 +1,101 @@
|
||||
import tempfile
|
||||
from pathlib import Path
|
||||
|
||||
from flask import Flask
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.api import graph_bp
|
||||
from app.api.auth import auth_bp
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.project import ProjectManager
|
||||
from app.models.saas import Organization
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
|
||||
|
||||
def make_graph_app():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SECRET_KEY="test-secret", SESSION_COOKIE_SECURE=False)
|
||||
app.extensions["crowdsight_session_factory"] = session_factory
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
app.register_blueprint(graph_bp, url_prefix="/api/graph")
|
||||
|
||||
with session_factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="org-a")
|
||||
user_a = repo.create_user(
|
||||
email="user-a@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(user_a.id, org.id, "user")
|
||||
user_b = repo.create_user(
|
||||
email="user-b@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(user_b.id, org.id, "user")
|
||||
admin = repo.create_user(
|
||||
email="admin@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(admin.id, org.id, "admin")
|
||||
session.commit()
|
||||
|
||||
return app, engine, user_a.id, user_b.id
|
||||
|
||||
|
||||
def login(client, email):
|
||||
response = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": email, "password": "correct horse battery staple"},
|
||||
)
|
||||
assert response.status_code == 200
|
||||
|
||||
|
||||
def test_graph_data_and_task_routes_require_auth():
|
||||
original_dir = ProjectManager.PROJECTS_DIR
|
||||
app, engine, _user_a_id, _user_b_id = make_graph_app()
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
ProjectManager.PROJECTS_DIR = str(Path(temp_dir) / "projects")
|
||||
try:
|
||||
client = app.test_client()
|
||||
assert client.get("/api/graph/data/graph-any").status_code == 401
|
||||
login(client, "user-a@example.com")
|
||||
assert client.get("/api/graph/data/arbitrary-graph").status_code == 404
|
||||
assert client.get("/api/graph/task/arbitrary-task").status_code == 404
|
||||
finally:
|
||||
ProjectManager.PROJECTS_DIR = original_dir
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_graph_project_routes_require_auth_and_scope_records():
|
||||
original_dir = ProjectManager.PROJECTS_DIR
|
||||
app, engine, user_a_id, user_b_id = make_graph_app()
|
||||
with tempfile.TemporaryDirectory() as temp_dir:
|
||||
ProjectManager.PROJECTS_DIR = str(Path(temp_dir) / "projects")
|
||||
try:
|
||||
project_a = ProjectManager.create_project(
|
||||
"A", organization_id="org_a_placeholder", owner_user_id=user_a_id
|
||||
)
|
||||
project_b = ProjectManager.create_project(
|
||||
"B", organization_id="org_a_placeholder", owner_user_id=user_b_id
|
||||
)
|
||||
# The app's organization id is discovered from the seeded user.
|
||||
with app.extensions["crowdsight_session_factory"]() as session:
|
||||
org_id = session.query(Organization).one().id
|
||||
for project in (project_a, project_b):
|
||||
project.organization_id = org_id
|
||||
ProjectManager.save_project(project)
|
||||
|
||||
client = app.test_client()
|
||||
assert client.get(f"/api/graph/project/{project_a.project_id}").status_code == 401
|
||||
|
||||
login(client, "user-a@example.com")
|
||||
assert client.get(f"/api/graph/project/{project_a.project_id}").status_code == 200
|
||||
assert client.get(f"/api/graph/project/{project_b.project_id}").status_code == 404
|
||||
listed = client.get("/api/graph/project/list").get_json()
|
||||
assert [item["project_id"] for item in listed["data"]] == [project_a.project_id]
|
||||
finally:
|
||||
ProjectManager.PROJECTS_DIR = original_dir
|
||||
engine.dispose()
|
||||
201
backend/tests/test_graph_backend_routes.py
Normal file
201
backend/tests/test_graph_backend_routes.py
Normal file
@@ -0,0 +1,201 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
import pytest
|
||||
from flask import Flask, g
|
||||
|
||||
import app.api.graph as graph_api
|
||||
from app.api import graph_bp
|
||||
from app.config import Config
|
||||
from app.models.project import ProjectManager
|
||||
from app.security.policy import Role
|
||||
from app.services.local_graph_builder import LocalGraphBuilderService
|
||||
|
||||
|
||||
class _LocalGraphSpy:
|
||||
instances = []
|
||||
|
||||
def __init__(self, session_factory, **kwargs):
|
||||
self.session_factory = session_factory
|
||||
self.kwargs = kwargs
|
||||
self.calls = []
|
||||
self.__class__.instances.append(self)
|
||||
|
||||
def get_graph_data(self, graph_id):
|
||||
self.calls.append(("get_graph_data", graph_id))
|
||||
return {"graph_id": graph_id, "node_count": 0, "edge_count": 0, "nodes": [], "edges": []}
|
||||
|
||||
def delete_graph(self, graph_id):
|
||||
self.calls.append(("delete_graph", graph_id))
|
||||
|
||||
|
||||
def make_app():
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SECRET_KEY="unit-test")
|
||||
app.extensions["crowdsight_session_factory"] = lambda: None
|
||||
app.register_blueprint(graph_bp, url_prefix="/api/graph")
|
||||
return app
|
||||
|
||||
|
||||
def set_actor():
|
||||
g.auth_context = SimpleNamespace(
|
||||
user=SimpleNamespace(id="user-a"),
|
||||
membership=SimpleNamespace(role=Role.USER),
|
||||
organization=SimpleNamespace(id="org-a"),
|
||||
)
|
||||
|
||||
|
||||
def test_delete_project_rejects_out_of_scope_project_before_delete(monkeypatch):
|
||||
app = make_app()
|
||||
delete_calls = []
|
||||
monkeypatch.setattr(graph_api, "_scoped_project", lambda project_id: None)
|
||||
monkeypatch.setattr(
|
||||
ProjectManager,
|
||||
"delete_project",
|
||||
lambda project_id, **kwargs: delete_calls.append((project_id, kwargs)) or True,
|
||||
)
|
||||
|
||||
with app.test_request_context("/api/graph/project/project-b", method="DELETE"):
|
||||
set_actor()
|
||||
result = graph_api.delete_project.__wrapped__("project-b")
|
||||
|
||||
if isinstance(result, tuple):
|
||||
response, status = result
|
||||
else:
|
||||
response, status = result, 200
|
||||
assert status == 404
|
||||
assert response.get_json()["success"] is False
|
||||
assert delete_calls == []
|
||||
|
||||
|
||||
def test_local_graph_data_route_uses_local_adapter_without_constructing_zep(monkeypatch):
|
||||
app = make_app()
|
||||
old_backend = Config.MEMORY_BACKEND
|
||||
old_key = Config.ZEP_API_KEY
|
||||
try:
|
||||
monkeypatch.setattr(Config, "MEMORY_BACKEND", "local")
|
||||
monkeypatch.setattr(Config, "ZEP_API_KEY", None)
|
||||
monkeypatch.setattr(
|
||||
graph_api,
|
||||
"_scoped_graph",
|
||||
lambda graph_id: SimpleNamespace(project_id="project-local"),
|
||||
)
|
||||
monkeypatch.setattr(graph_api, "LocalGraphBuilderService", _LocalGraphSpy)
|
||||
monkeypatch.setattr(
|
||||
graph_api,
|
||||
"GraphBuilderService",
|
||||
lambda **kwargs: pytest.fail("local graph read must not construct Zep"),
|
||||
)
|
||||
_LocalGraphSpy.instances.clear()
|
||||
|
||||
with app.test_request_context("/api/graph/data/graph-local"):
|
||||
set_actor()
|
||||
response = graph_api.get_graph_data.__wrapped__("graph-local")
|
||||
|
||||
if isinstance(response, tuple):
|
||||
assert response[1] == 200
|
||||
payload = response[0].get_json()
|
||||
else:
|
||||
payload = response.get_json()
|
||||
assert payload["success"] is True
|
||||
assert payload["data"]["graph_id"] == "graph-local"
|
||||
assert _LocalGraphSpy.instances[0].kwargs["organization_id"] == "org-a"
|
||||
assert _LocalGraphSpy.instances[0].kwargs["project_id"] == "project-local"
|
||||
assert _LocalGraphSpy.instances[0].calls == [("get_graph_data", "graph-local")]
|
||||
finally:
|
||||
Config.MEMORY_BACKEND = old_backend
|
||||
Config.ZEP_API_KEY = old_key
|
||||
|
||||
|
||||
def test_local_graph_delete_route_uses_local_adapter_without_constructing_zep(monkeypatch):
|
||||
app = make_app()
|
||||
old_backend = Config.MEMORY_BACKEND
|
||||
old_key = Config.ZEP_API_KEY
|
||||
try:
|
||||
monkeypatch.setattr(Config, "MEMORY_BACKEND", "local")
|
||||
monkeypatch.setattr(Config, "ZEP_API_KEY", None)
|
||||
monkeypatch.setattr(
|
||||
graph_api,
|
||||
"_scoped_graph",
|
||||
lambda graph_id: SimpleNamespace(project_id="project-local"),
|
||||
)
|
||||
monkeypatch.setattr(graph_api, "LocalGraphBuilderService", _LocalGraphSpy)
|
||||
monkeypatch.setattr(
|
||||
graph_api,
|
||||
"GraphBuilderService",
|
||||
lambda **kwargs: pytest.fail("local graph delete must not construct Zep"),
|
||||
)
|
||||
_LocalGraphSpy.instances.clear()
|
||||
|
||||
with app.test_request_context("/api/graph/delete/graph-local", method="DELETE"):
|
||||
set_actor()
|
||||
response = graph_api.delete_graph.__wrapped__("graph-local")
|
||||
|
||||
payload = response[0].get_json() if isinstance(response, tuple) else response.get_json()
|
||||
assert payload["success"] is True
|
||||
assert _LocalGraphSpy.instances[0].calls == [("delete_graph", "graph-local")]
|
||||
finally:
|
||||
Config.MEMORY_BACKEND = old_backend
|
||||
Config.ZEP_API_KEY = old_key
|
||||
|
||||
|
||||
def test_local_graph_builder_deletes_only_the_scoped_graph():
|
||||
from sqlalchemy import create_engine
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.memory import MemoryEdge, MemoryGraph, MemoryNode
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
try:
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_client=object(),
|
||||
)
|
||||
graph_id = builder.create_graph(name="Local")
|
||||
|
||||
with session_factory() as session:
|
||||
first = MemoryNode(
|
||||
id="node-a",
|
||||
graph_id=graph_id,
|
||||
canonical_name="A",
|
||||
normalized_name="a",
|
||||
)
|
||||
second = MemoryNode(
|
||||
id="node-b",
|
||||
graph_id=graph_id,
|
||||
canonical_name="B",
|
||||
normalized_name="b",
|
||||
)
|
||||
session.add_all([first, second])
|
||||
session.flush()
|
||||
session.add(
|
||||
MemoryEdge(
|
||||
id="edge-ab",
|
||||
graph_id=graph_id,
|
||||
source_node_id=first.id,
|
||||
target_node_id=second.id,
|
||||
relation="RELATED",
|
||||
fact="A is related to B.",
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
other_builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-b",
|
||||
project_id="project-b",
|
||||
extraction_client=object(),
|
||||
)
|
||||
with pytest.raises(ValueError, match="memory_graph_not_found"):
|
||||
other_builder.delete_graph(graph_id)
|
||||
|
||||
builder.delete_graph(graph_id)
|
||||
with session_factory() as session:
|
||||
assert session.get(MemoryGraph, graph_id) is None
|
||||
assert session.get(MemoryNode, "node-a") is None
|
||||
assert session.get(MemoryNode, "node-b") is None
|
||||
assert session.get(MemoryEdge, "edge-ab") is None
|
||||
finally:
|
||||
engine.dispose()
|
||||
49
backend/tests/test_graph_builder_task_scope.py
Normal file
49
backend/tests/test_graph_builder_task_scope.py
Normal file
@@ -0,0 +1,49 @@
|
||||
from app.db import Base, create_database_engine, create_session_factory
|
||||
from app.models.task import TaskManager
|
||||
from app.services import graph_builder
|
||||
from app.services.identity import IdentityRepository
|
||||
|
||||
|
||||
def test_graph_builder_async_task_keeps_organization_scope(monkeypatch):
|
||||
engine = create_database_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
|
||||
class FakeZep:
|
||||
def __init__(self, api_key):
|
||||
self.api_key = api_key
|
||||
|
||||
class FakeThread:
|
||||
def __init__(self, *, target, args):
|
||||
self.target = target
|
||||
self.args = args
|
||||
self.daemon = False
|
||||
|
||||
def start(self):
|
||||
return None
|
||||
|
||||
monkeypatch.setattr(graph_builder, "Zep", FakeZep)
|
||||
monkeypatch.setattr(graph_builder.threading, "Thread", FakeThread)
|
||||
|
||||
try:
|
||||
with factory() as session:
|
||||
organization = IdentityRepository(session).create_organization(
|
||||
name="Graph Builder Tests", slug="graph-builder-tests"
|
||||
)
|
||||
session.commit()
|
||||
|
||||
builder = graph_builder.GraphBuilderService(
|
||||
api_key="unit-test-placeholder",
|
||||
organization_id=organization.id,
|
||||
session_factory=factory,
|
||||
)
|
||||
task_id = builder.build_graph_async("seed text", {"entity_types": []})
|
||||
|
||||
task = TaskManager(session_factory=factory).get_task(
|
||||
task_id, organization_id=organization.id
|
||||
)
|
||||
assert task is not None
|
||||
assert task.metadata["organization_id"] == organization.id
|
||||
finally:
|
||||
TaskManager.configure(None)
|
||||
engine.dispose()
|
||||
184
backend/tests/test_idempotency_api.py
Normal file
184
backend/tests/test_idempotency_api.py
Normal file
@@ -0,0 +1,184 @@
|
||||
import secrets
|
||||
|
||||
from flask import Flask, jsonify
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.security.auth import require_auth
|
||||
from app.security.policy import Role
|
||||
from app.services.idempotency import IdempotencyService, _request_fingerprint_payload, idempotent
|
||||
from app.services.identity import IdentityRepository, PasswordService
|
||||
from app.utils.api_errors import ApiError
|
||||
from app.utils.locale import t
|
||||
|
||||
|
||||
def make_app():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
app = Flask(__name__)
|
||||
app.config.update(TESTING=True, SESSION_COOKIE_SECURE=False)
|
||||
app.config["SECRET_KEY"] = secrets.token_hex(32)
|
||||
app.extensions["crowdsight_session_factory"] = session_factory
|
||||
|
||||
@app.errorhandler(ApiError)
|
||||
def handle_api_error(error):
|
||||
return {**error.to_payload(t)}, error.status_code
|
||||
|
||||
@app.post("/mutate")
|
||||
@require_auth
|
||||
@idempotent
|
||||
def mutate():
|
||||
return jsonify({"success": True, "data": {"value": "created"}}), 201
|
||||
|
||||
@app.post("/mutate-alt")
|
||||
@require_auth
|
||||
@idempotent
|
||||
def mutate_alt():
|
||||
return jsonify({"success": True, "data": {"value": "alternate"}}), 201
|
||||
|
||||
with session_factory() as session:
|
||||
repo = IdentityRepository(session)
|
||||
org = repo.create_organization(name="Org A", slug="idempotent-org")
|
||||
user = repo.create_user(
|
||||
email="idempotent@example.com",
|
||||
password_hash=PasswordService.hash_password("correct horse battery staple"),
|
||||
)
|
||||
repo.create_membership(user.id, org.id, Role.ADMIN)
|
||||
session.commit()
|
||||
return app, engine
|
||||
|
||||
|
||||
def login(client):
|
||||
response = client.post(
|
||||
"/mutate",
|
||||
json={"value": "created"},
|
||||
headers={"Idempotency-Key": "mutation-1"},
|
||||
)
|
||||
return response
|
||||
|
||||
|
||||
def test_idempotent_route_replays_completed_response():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
auth_app = Flask(__name__)
|
||||
# Login through the real blueprint in a small shared app is covered separately;
|
||||
# seed the opaque cookie via the auth endpoint mounted on this test app.
|
||||
from app.api.auth import auth_bp
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
login_response = client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": "idempotent@example.com", "password": "correct horse battery staple"},
|
||||
)
|
||||
assert login_response.status_code == 200
|
||||
csrf = client.get_cookie("crowdsight_csrf").value
|
||||
headers = {"Idempotency-Key": "mutation-1", "X-CSRF-Token": csrf}
|
||||
|
||||
first = client.post("/mutate", json={"value": "created"}, headers=headers)
|
||||
second = client.post("/mutate", json={"value": "created"}, headers=headers)
|
||||
assert first.status_code == second.status_code == 201
|
||||
assert first.get_json() == second.get_json()
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_idempotent_route_rejects_same_key_for_different_body():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
from app.api.auth import auth_bp
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
assert client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": "idempotent@example.com", "password": "correct horse battery staple"},
|
||||
).status_code == 200
|
||||
csrf = client.get_cookie("crowdsight_csrf").value
|
||||
headers = {"Idempotency-Key": "mutation-2", "X-CSRF-Token": csrf}
|
||||
assert client.post("/mutate", json={"value": "a"}, headers=headers).status_code == 201
|
||||
conflict = client.post("/mutate", json={"value": "b"}, headers=headers)
|
||||
assert conflict.status_code == 409
|
||||
assert conflict.get_json()["error_code"] == "idempotency_key_reused"
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_multipart_file_content_affects_fingerprint():
|
||||
app = Flask(__name__)
|
||||
boundary = b"FixedBoundary"
|
||||
|
||||
def fingerprint(content):
|
||||
body = (
|
||||
b"--" + boundary + b"\r\n"
|
||||
b'Content-Disposition: form-data; name="simulation_requirement"\r\n\r\n'
|
||||
b"req\r\n"
|
||||
b"--" + boundary + b"\r\n"
|
||||
b'Content-Disposition: form-data; name="files"; filename="doc.txt"\r\n'
|
||||
b"Content-Type: text/plain\r\n\r\n"
|
||||
+ content
|
||||
+ b"\r\n--"
|
||||
+ boundary
|
||||
+ b"--\r\n"
|
||||
)
|
||||
with app.test_request_context(
|
||||
"/upload",
|
||||
method="POST",
|
||||
data=body,
|
||||
content_type="multipart/form-data; boundary=FixedBoundary",
|
||||
):
|
||||
payload = _request_fingerprint_payload()
|
||||
return IdempotencyService._request_hash(payload)
|
||||
|
||||
assert fingerprint(b"alpha") != fingerprint(b"bravo")
|
||||
|
||||
|
||||
def test_all_multipart_files_affect_fingerprint():
|
||||
app = Flask(__name__)
|
||||
boundary = b"FixedBoundary"
|
||||
|
||||
def fingerprint(second_content):
|
||||
first_part = (
|
||||
b'Content-Disposition: form-data; name="files"; filename="first.txt"\r\n'
|
||||
b"Content-Type: text/plain\r\n\r\nalpha"
|
||||
)
|
||||
second_part = (
|
||||
b'Content-Disposition: form-data; name="files"; filename="second.txt"\r\n'
|
||||
b"Content-Type: text/plain\r\n\r\n"
|
||||
+ second_content
|
||||
)
|
||||
body = (
|
||||
b"--" + boundary + b"\r\n"
|
||||
+ first_part
|
||||
+ b"\r\n--" + boundary + b"\r\n"
|
||||
+ second_part
|
||||
+ b"\r\n--" + boundary + b"--\r\n"
|
||||
)
|
||||
with app.test_request_context(
|
||||
"/upload",
|
||||
method="POST",
|
||||
data=body,
|
||||
content_type="multipart/form-data; boundary=FixedBoundary",
|
||||
):
|
||||
return IdempotencyService._request_hash(_request_fingerprint_payload())
|
||||
|
||||
assert fingerprint(b"bravo") != fingerprint(b"charl")
|
||||
|
||||
|
||||
def test_idempotency_key_cannot_replay_across_routes():
|
||||
app, engine = make_app()
|
||||
try:
|
||||
client = app.test_client()
|
||||
from app.api.auth import auth_bp
|
||||
app.register_blueprint(auth_bp, url_prefix="/api/auth")
|
||||
assert client.post(
|
||||
"/api/auth/login",
|
||||
json={"email": "idempotent@example.com", "password": "correct horse battery staple"},
|
||||
).status_code == 200
|
||||
csrf = client.get_cookie("crowdsight_csrf").value
|
||||
headers = {"Idempotency-Key": "cross-route", "X-CSRF-Token": csrf}
|
||||
assert client.post("/mutate", json={"value": "created"}, headers=headers).status_code == 201
|
||||
conflict = client.post("/mutate-alt", json={"value": "created"}, headers=headers)
|
||||
assert conflict.status_code == 409
|
||||
assert conflict.get_json()["error_code"] == "idempotency_key_reused"
|
||||
finally:
|
||||
engine.dispose()
|
||||
161
backend/tests/test_job_queue.py
Normal file
161
backend/tests/test_job_queue.py
Normal file
@@ -0,0 +1,161 @@
|
||||
"""TDD gate: durable job queue claiming/consumption without a broker.
|
||||
|
||||
A real worker later runs on a queue provider, but the claim/complete lifecycle
|
||||
and tenant scope must work against the durable ``jobs`` table now so jobs
|
||||
survive restarts and are isolated per organization.
|
||||
"""
|
||||
|
||||
import pytest
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.operations import Job, JobStatus
|
||||
from app.services.job_queue import JobQueue
|
||||
|
||||
|
||||
@pytest.fixture()
|
||||
def session_factory():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
try:
|
||||
yield factory
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def _seed_job(session_factory, *, organization_id="org-a", operation="graph.build"):
|
||||
session = session_factory()
|
||||
job = Job(
|
||||
organization_id=organization_id,
|
||||
owner_user_id="user-1",
|
||||
operation=operation,
|
||||
status=JobStatus.QUEUED,
|
||||
)
|
||||
session.add(job)
|
||||
session.commit()
|
||||
job_id = job.id
|
||||
session.close()
|
||||
return job_id
|
||||
|
||||
|
||||
def test_job_queue_claims_single_queued_job(session_factory):
|
||||
job_id = _seed_job(session_factory)
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
claimed = queue.claim_next_job(worker_id="worker-1")
|
||||
assert claimed is not None
|
||||
assert claimed.id == job_id
|
||||
assert claimed.status == JobStatus.RUNNING.value
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_job_queue_does_not_double_claim(session_factory):
|
||||
_seed_job(session_factory)
|
||||
|
||||
s1 = session_factory()
|
||||
s2 = session_factory()
|
||||
try:
|
||||
q1 = JobQueue(s1)
|
||||
q2 = JobQueue(s2)
|
||||
first = q1.claim_next_job(worker_id="w-1")
|
||||
second = q2.claim_next_job(worker_id="w-2")
|
||||
assert first is not None
|
||||
# Second worker must not see the already-claimed job (and no other jobs).
|
||||
assert second is None
|
||||
finally:
|
||||
s1.close()
|
||||
s2.close()
|
||||
|
||||
|
||||
def test_job_queue_does_not_claim_other_tenants_job(session_factory):
|
||||
_seed_job(session_factory, organization_id="org-a")
|
||||
_seed_job(session_factory, organization_id="org-b")
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
# A worker processing org-a claims only org-a jobs.
|
||||
claimed = queue.claim_next_job(worker_id="w-1", organization_id="org-a")
|
||||
assert claimed is not None and claimed.organization_id == "org-a"
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_job_queue_complete_and_fail_update_status(session_factory):
|
||||
job_id = _seed_job(session_factory)
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
claimed = queue.claim_next_job(worker_id="w-1")
|
||||
assert claimed is not None
|
||||
queue.complete_job(claimed.id, result={"ok": True})
|
||||
session.commit()
|
||||
refreshed = session.get(Job, claimed.id)
|
||||
assert refreshed.status == JobStatus.SUCCEEDED.value
|
||||
assert refreshed.result == {"ok": True}
|
||||
|
||||
job_id2 = _seed_job(session_factory)
|
||||
claimed2 = queue.claim_next_job(worker_id="w-1")
|
||||
assert claimed2 is not None and claimed2.id == job_id2
|
||||
queue.fail_job(claimed2.id, error_code="boom")
|
||||
session.commit()
|
||||
refreshed2 = session.get(Job, claimed2.id)
|
||||
assert refreshed2.status == JobStatus.FAILED.value
|
||||
assert refreshed2.error_code == "boom"
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_job_queue_none_when_empty(session_factory):
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
assert queue.claim_next_job(worker_id="w-1") is None
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_job_queue_dispatch_invokes_registered_handler(session_factory):
|
||||
calls = {}
|
||||
|
||||
def handler(payload, job):
|
||||
calls["payload"] = payload
|
||||
calls["job_id"] = job.id
|
||||
return {"processed": True}
|
||||
|
||||
job_id = _seed_job(session_factory, operation="graph.build")
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
queue.register_handler("graph.build", handler)
|
||||
claimed = queue.claim_next_job(worker_id="w-1")
|
||||
assert claimed is not None
|
||||
result = queue.dispatch(claimed, payload={"graph_id": "g1"})
|
||||
assert result == {"processed": True}
|
||||
assert calls["job_id"] == job_id
|
||||
assert calls["payload"] == {"graph_id": "g1"}
|
||||
finally:
|
||||
session.close()
|
||||
|
||||
|
||||
def test_job_queue_dispatch_fails_unhandled_operation(session_factory):
|
||||
job_id = _seed_job(session_factory, operation="unknown.op")
|
||||
|
||||
session = session_factory()
|
||||
try:
|
||||
queue = JobQueue(session)
|
||||
claimed = queue.claim_next_job(worker_id="w-1")
|
||||
assert claimed is not None
|
||||
with pytest.raises(ValueError, match="no_handler"):
|
||||
queue.dispatch(claimed, payload={})
|
||||
# Claimed job is still running until the worker decides to fail it.
|
||||
refreshed = session.get(Job, claimed.id)
|
||||
assert refreshed.status == JobStatus.RUNNING.value
|
||||
finally:
|
||||
session.close()
|
||||
47
backend/tests/test_language_policy.py
Normal file
47
backend/tests/test_language_policy.py
Normal file
@@ -0,0 +1,47 @@
|
||||
import importlib.util
|
||||
from pathlib import Path
|
||||
import sys
|
||||
import unittest
|
||||
|
||||
|
||||
_MODULE_PATH = Path(__file__).parents[1] / "app" / "utils" / "language_policy.py"
|
||||
_SPEC = importlib.util.spec_from_file_location("language_policy_under_test", _MODULE_PATH)
|
||||
assert _SPEC is not None and _SPEC.loader is not None
|
||||
_LANGUAGE_POLICY = importlib.util.module_from_spec(_SPEC)
|
||||
sys.modules[_SPEC.name] = _LANGUAGE_POLICY
|
||||
_SPEC.loader.exec_module(_LANGUAGE_POLICY)
|
||||
|
||||
DEFAULT_LOCALE = _LANGUAGE_POLICY.DEFAULT_LOCALE
|
||||
SUPPORTED_LOCALES = _LANGUAGE_POLICY.SUPPORTED_LOCALES
|
||||
locale_from_accept_language = _LANGUAGE_POLICY.locale_from_accept_language
|
||||
normalize_locale = _LANGUAGE_POLICY.normalize_locale
|
||||
|
||||
|
||||
class LanguagePolicyTests(unittest.TestCase):
|
||||
def test_supported_locales_are_thai_and_english(self):
|
||||
self.assertEqual(SUPPORTED_LOCALES, ("th", "en"))
|
||||
self.assertEqual(DEFAULT_LOCALE, "th")
|
||||
|
||||
def test_normalize_locale_accepts_language_region_values(self):
|
||||
self.assertEqual(normalize_locale("en-US"), "en")
|
||||
self.assertEqual(normalize_locale("TH_th"), "th")
|
||||
|
||||
def test_normalize_locale_migrates_legacy_chinese_to_default(self):
|
||||
self.assertEqual(normalize_locale("zh"), "th")
|
||||
self.assertEqual(normalize_locale("zh-CN"), "th")
|
||||
|
||||
def test_normalize_locale_falls_back_for_unknown_values(self):
|
||||
self.assertEqual(normalize_locale(None), "th")
|
||||
self.assertEqual(normalize_locale("fr"), "th")
|
||||
self.assertEqual(normalize_locale(""), "th")
|
||||
|
||||
def test_accept_language_uses_quality_and_skips_zero_quality(self):
|
||||
header = "zh-CN;q=1, fr-FR;q=0.9, en-US;q=0.8, th;q=0"
|
||||
self.assertEqual(locale_from_accept_language(header), "en")
|
||||
|
||||
def test_accept_language_returns_default_when_no_supported_language_exists(self):
|
||||
self.assertEqual(locale_from_accept_language("ja-JP, de;q=0.8"), "th")
|
||||
|
||||
|
||||
if __name__ == "__main__":
|
||||
unittest.main()
|
||||
239
backend/tests/test_local_golden_flow.py
Normal file
239
backend/tests/test_local_golden_flow.py
Normal file
@@ -0,0 +1,239 @@
|
||||
from types import SimpleNamespace
|
||||
from typing import Any, cast
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.config import Config
|
||||
from app.db import Base, create_session_factory
|
||||
from app.services.local_graph_builder import LocalGraphBuilderService
|
||||
from app.services.memory_entity_reader import LocalEntityReader
|
||||
from app.services.memory_extraction import (
|
||||
ExtractedEdge,
|
||||
ExtractedEntity,
|
||||
MemoryExtractionResult,
|
||||
)
|
||||
from app.services.memory_service import MemoryExtractionService
|
||||
from app.services.memory_tools import LocalMemoryTools
|
||||
from app.services.oasis_profile_generator import OasisProfileGenerator
|
||||
from app.services.report_agent import ReportAgent
|
||||
from app.services.report_agent import ReportManager, ReportOutline, ReportSection, ReportStatus
|
||||
from app.services.simulation_config_generator import (
|
||||
AgentActivityConfig,
|
||||
SimulationConfigGenerator,
|
||||
SimulationParameters,
|
||||
)
|
||||
from app.services.simulation_manager import SimulationManager, SimulationStatus
|
||||
from app.services.memory_entity_reader import make_local_entity_reader_factory
|
||||
|
||||
|
||||
class DeterministicExtraction:
|
||||
def extract(self, *, language, ontology, episode_text, context=""):
|
||||
return MemoryExtractionResult(
|
||||
entities=[
|
||||
ExtractedEntity(
|
||||
mention="Alice",
|
||||
canonical_name="Alice",
|
||||
labels=["Entity", "Person"],
|
||||
summary="A founder building Orbit.",
|
||||
confidence=0.99,
|
||||
),
|
||||
ExtractedEntity(
|
||||
mention="Orbit",
|
||||
canonical_name="Orbit",
|
||||
labels=["Entity", "Organization"],
|
||||
summary="A local project.",
|
||||
confidence=0.99,
|
||||
),
|
||||
],
|
||||
edges=[
|
||||
ExtractedEdge(
|
||||
source_entity_ref="Alice",
|
||||
target_entity_ref="Orbit",
|
||||
relation="FOUNDED",
|
||||
fact="Alice founded Orbit.",
|
||||
confidence=0.98,
|
||||
)
|
||||
],
|
||||
episode_summary="Alice founded Orbit.",
|
||||
)
|
||||
|
||||
def persist(self, repository, result, *, source_type, source_ref, episode_text):
|
||||
return MemoryExtractionService(cast(Any, None)).persist(
|
||||
repository,
|
||||
result,
|
||||
source_type=source_type,
|
||||
source_ref=source_ref,
|
||||
episode_text=episode_text,
|
||||
)
|
||||
|
||||
|
||||
def test_local_graph_profile_report_golden_flow(monkeypatch):
|
||||
monkeypatch.setattr(Config, "MEMORY_BACKEND", "local")
|
||||
monkeypatch.setattr(Config, "LLM_API_KEY", "test-key")
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
try:
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_service=cast(Any, DeterministicExtraction()),
|
||||
language="en",
|
||||
)
|
||||
graph_id = builder.create_graph("golden")
|
||||
builder.add_text_batches(graph_id, ["Alice founded Orbit."])
|
||||
|
||||
reader = LocalEntityReader(
|
||||
session_factory(),
|
||||
organization_id="org-a",
|
||||
graph_id=graph_id,
|
||||
owns_session=True,
|
||||
)
|
||||
try:
|
||||
filtered = reader.filter_defined_entities(defined_entity_types=["Person"])
|
||||
assert [entity.name for entity in filtered.entities] == ["Alice"]
|
||||
|
||||
local_tools = LocalMemoryTools(reader.repository)
|
||||
profile_generator = OasisProfileGenerator(
|
||||
graph_id=graph_id,
|
||||
use_zep_context=True,
|
||||
local_memory_tools=local_tools,
|
||||
)
|
||||
profile = profile_generator.generate_profile_from_entity(
|
||||
cast(Any, filtered.entities[0]),
|
||||
user_id=1,
|
||||
use_llm=False,
|
||||
)
|
||||
assert profile.name == "Alice"
|
||||
assert profile.source_entity_uuid == filtered.entities[0].uuid
|
||||
assert profile_generator.zep_client is None
|
||||
|
||||
report_agent = ReportAgent(
|
||||
graph_id=graph_id,
|
||||
simulation_id="simulation-a",
|
||||
simulation_requirement="Understand the project origin.",
|
||||
llm_client=cast(Any, SimpleNamespace()),
|
||||
memory_tools=local_tools,
|
||||
)
|
||||
report_context = report_agent._execute_tool(
|
||||
"quick_search",
|
||||
{"query": "Alice", "limit": 10},
|
||||
)
|
||||
assert "Alice founded Orbit." in report_context
|
||||
finally:
|
||||
reader.close()
|
||||
finally:
|
||||
engine.dispose()
|
||||
|
||||
|
||||
def test_local_graph_profile_simulation_report_is_persisted(monkeypatch, tmp_path):
|
||||
monkeypatch.setattr(Config, "MEMORY_BACKEND", "local")
|
||||
monkeypatch.setattr(Config, "LLM_API_KEY", "test-key")
|
||||
monkeypatch.setattr(Config, "UPLOAD_FOLDER", str(tmp_path / "uploads"))
|
||||
monkeypatch.setattr(ReportManager, "REPORTS_DIR", str(tmp_path / "reports"))
|
||||
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
simulation_dir = tmp_path / "simulations"
|
||||
try:
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_service=cast(Any, DeterministicExtraction()),
|
||||
language="en",
|
||||
)
|
||||
graph_id = builder.create_graph("e2e")
|
||||
builder.add_text_batches(graph_id, ["Alice founded Orbit."])
|
||||
|
||||
monkeypatch.setattr(
|
||||
SimulationConfigGenerator,
|
||||
"generate_config",
|
||||
lambda self, **kwargs: SimulationParameters(
|
||||
simulation_id=kwargs["simulation_id"],
|
||||
project_id=kwargs["project_id"],
|
||||
graph_id=kwargs["graph_id"],
|
||||
simulation_requirement=kwargs["simulation_requirement"],
|
||||
agent_configs=[
|
||||
AgentActivityConfig(
|
||||
agent_id=1,
|
||||
entity_uuid="node-alice",
|
||||
entity_name="Alice",
|
||||
entity_type="Person",
|
||||
)
|
||||
],
|
||||
generation_reasoning="deterministic-test",
|
||||
),
|
||||
)
|
||||
|
||||
simulation_manager = SimulationManager(
|
||||
entity_reader_factory=make_local_entity_reader_factory(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
)
|
||||
)
|
||||
simulation_manager.SIMULATION_DATA_DIR = str(simulation_dir)
|
||||
simulation = simulation_manager.create_simulation(
|
||||
project_id="project-a",
|
||||
graph_id=graph_id,
|
||||
enable_twitter=False,
|
||||
enable_reddit=True,
|
||||
)
|
||||
prepared = simulation_manager.prepare_simulation(
|
||||
simulation.simulation_id,
|
||||
simulation_requirement="Understand the project origin.",
|
||||
document_text="Alice founded Orbit.",
|
||||
defined_entity_types=["Person"],
|
||||
use_llm_for_profiles=False,
|
||||
parallel_profile_count=1,
|
||||
)
|
||||
assert prepared.status is SimulationStatus.READY
|
||||
assert prepared.entities_count == 1
|
||||
assert prepared.profiles_count == 1
|
||||
assert simulation_manager.get_profiles(prepared.simulation_id) [0]["name"] == "Alice"
|
||||
|
||||
report_session = session_factory()
|
||||
try:
|
||||
report_tools = LocalMemoryTools(
|
||||
report_session,
|
||||
organization_id="org-a",
|
||||
graph_id=graph_id,
|
||||
)
|
||||
report_agent = ReportAgent(
|
||||
graph_id=graph_id,
|
||||
simulation_id=prepared.simulation_id,
|
||||
simulation_requirement="Understand the project origin.",
|
||||
llm_client=cast(Any, SimpleNamespace()),
|
||||
memory_tools=report_tools,
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ReportAgent,
|
||||
"plan_outline",
|
||||
lambda self, progress_callback=None: ReportOutline(
|
||||
title="Local E2E Report",
|
||||
summary="Evidence-backed local report.",
|
||||
sections=[ReportSection(title="Evidence")],
|
||||
),
|
||||
)
|
||||
monkeypatch.setattr(
|
||||
ReportAgent,
|
||||
"_generate_section_react",
|
||||
lambda self, section, outline, previous_sections, progress_callback=None, section_index=0: self._execute_tool(
|
||||
"quick_search",
|
||||
{"query": "Alice", "limit": 10},
|
||||
),
|
||||
)
|
||||
|
||||
report = report_agent.generate_report(report_id="report-local-e2e")
|
||||
assert report.status is ReportStatus.COMPLETED
|
||||
assert "Alice founded Orbit." in report.markdown_content
|
||||
persisted = ReportManager.get_report("report-local-e2e")
|
||||
assert persisted is not None
|
||||
assert persisted.status is ReportStatus.COMPLETED
|
||||
assert "Alice founded Orbit." in persisted.markdown_content
|
||||
finally:
|
||||
report_session.close()
|
||||
finally:
|
||||
engine.dispose()
|
||||
134
backend/tests/test_local_graph_builder.py
Normal file
134
backend/tests/test_local_graph_builder.py
Normal file
@@ -0,0 +1,134 @@
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.memory import MemoryGraph
|
||||
from app.services.local_graph_builder import LocalGraphBuilderService
|
||||
|
||||
|
||||
class FakeExtractionClient:
|
||||
def __init__(self):
|
||||
self.calls = []
|
||||
|
||||
def chat_json(self, messages, temperature=0.3, max_tokens=4096):
|
||||
self.calls.append(messages)
|
||||
return {
|
||||
"entities": [
|
||||
{
|
||||
"mention": "Alice",
|
||||
"canonical_name": "Alice",
|
||||
"labels": ["Person"],
|
||||
"aliases": [],
|
||||
"attributes": {"role": "founder"},
|
||||
"summary": "Alice founded the local project.",
|
||||
"confidence": 0.95,
|
||||
},
|
||||
{
|
||||
"mention": "Orbit",
|
||||
"canonical_name": "Orbit",
|
||||
"labels": ["Project"],
|
||||
"aliases": [],
|
||||
"attributes": {},
|
||||
"summary": "Orbit is the project discussed in the episode.",
|
||||
"confidence": 0.9,
|
||||
},
|
||||
],
|
||||
"edges": [
|
||||
{
|
||||
"source_entity_ref": "Alice",
|
||||
"target_entity_ref": "Orbit",
|
||||
"relation": "FOUNDED",
|
||||
"fact": "Alice founded Orbit.",
|
||||
"attributes": {},
|
||||
"valid_at": None,
|
||||
"invalid_at": None,
|
||||
"expired_at": None,
|
||||
"confidence": 0.92,
|
||||
"evidence": ["episode-0"],
|
||||
}
|
||||
],
|
||||
"episode_summary": "Alice founded Orbit.",
|
||||
"unresolved_mentions": [],
|
||||
}
|
||||
|
||||
|
||||
def make_session_factory():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
return engine, create_session_factory(engine)
|
||||
|
||||
|
||||
def test_local_graph_builder_persists_scoped_graph_and_extracted_memory():
|
||||
engine, session_factory = make_session_factory()
|
||||
client = FakeExtractionClient()
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_client=client,
|
||||
)
|
||||
|
||||
graph_id = builder.create_graph(name="Orbit")
|
||||
builder.set_ontology(graph_id, {"entity_types": ["Person", "Project"], "relations": ["FOUNDED"]})
|
||||
episode_ids = builder.add_text_batches(
|
||||
graph_id,
|
||||
["Alice founded Orbit."],
|
||||
batch_size=3,
|
||||
)
|
||||
|
||||
assert len(episode_ids) == 1
|
||||
assert episode_ids[0].startswith("episode_")
|
||||
assert client.calls
|
||||
data = builder.get_graph_data(graph_id)
|
||||
assert data["graph_id"] == graph_id
|
||||
assert data["node_count"] == 2
|
||||
assert data["edge_count"] == 1
|
||||
|
||||
with session_factory() as session:
|
||||
graph = session.get(MemoryGraph, graph_id)
|
||||
assert graph.organization_id == "org-a"
|
||||
assert graph.project_id == "project-a"
|
||||
assert graph.ontology["relations"] == ["FOUNDED"]
|
||||
|
||||
|
||||
def test_local_graph_builder_reprocessing_is_idempotent_for_episode_and_edge():
|
||||
engine, session_factory = make_session_factory()
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_client=FakeExtractionClient(),
|
||||
)
|
||||
graph_id = builder.create_graph(name="Orbit")
|
||||
builder.set_ontology(graph_id, {"entity_types": ["Person", "Project"]})
|
||||
|
||||
builder.add_text_batches(graph_id, ["Alice founded Orbit."])
|
||||
builder.add_text_batches(graph_id, ["Alice founded Orbit."])
|
||||
|
||||
data = builder.get_graph_data(graph_id)
|
||||
assert data["node_count"] == 2
|
||||
assert data["edge_count"] == 1
|
||||
|
||||
|
||||
def test_local_graph_builder_fails_closed_for_wrong_organization():
|
||||
engine, session_factory = make_session_factory()
|
||||
builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
extraction_client=FakeExtractionClient(),
|
||||
)
|
||||
graph_id = builder.create_graph(name="Orbit")
|
||||
|
||||
other_builder = LocalGraphBuilderService(
|
||||
session_factory,
|
||||
organization_id="org-b",
|
||||
project_id="project-b",
|
||||
extraction_client=FakeExtractionClient(),
|
||||
)
|
||||
|
||||
try:
|
||||
other_builder.get_graph_data(graph_id)
|
||||
except ValueError as exc:
|
||||
assert str(exc) == "memory_graph_not_found"
|
||||
else:
|
||||
raise AssertionError("wrong organization must not read graph")
|
||||
50
backend/tests/test_local_graph_memory_updater.py
Normal file
50
backend/tests/test_local_graph_memory_updater.py
Normal file
@@ -0,0 +1,50 @@
|
||||
from datetime import datetime, timezone
|
||||
|
||||
from sqlalchemy import create_engine, select
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.memory import MemoryEpisode, MemoryGraph
|
||||
from app.services.local_graph_memory_updater import LocalGraphMemoryUpdater
|
||||
from app.services.memory_activity import AgentActivity
|
||||
|
||||
|
||||
def test_local_graph_memory_updater_persists_scoped_activity_episode():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
try:
|
||||
with session_factory() as session:
|
||||
session.add(
|
||||
MemoryGraph(
|
||||
id="graph-runtime",
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
updater = LocalGraphMemoryUpdater(
|
||||
simulation_id="sim-runtime",
|
||||
graph_id="graph-runtime",
|
||||
organization_id="org-a",
|
||||
session_factory=session_factory,
|
||||
)
|
||||
updater.start()
|
||||
updater.add_activity(
|
||||
AgentActivity(
|
||||
platform="twitter",
|
||||
agent_id=7,
|
||||
agent_name="Alice",
|
||||
action_type="CREATE_POST",
|
||||
action_args={"content": "Local memory is durable."},
|
||||
round_num=2,
|
||||
timestamp=datetime.now(timezone.utc).isoformat(),
|
||||
)
|
||||
)
|
||||
updater.stop()
|
||||
|
||||
with session_factory() as session:
|
||||
episodes = list(session.scalars(select(MemoryEpisode)))
|
||||
assert len(episodes) == 1
|
||||
finally:
|
||||
engine.dispose()
|
||||
31
backend/tests/test_local_reader_factory.py
Normal file
31
backend/tests/test_local_reader_factory.py
Normal file
@@ -0,0 +1,31 @@
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.memory import MemoryGraph
|
||||
from app.services.memory_entity_reader import LocalEntityReader, make_local_entity_reader_factory
|
||||
|
||||
|
||||
def test_local_reader_factory_scopes_each_reader_and_owns_worker_session():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
session_factory = create_session_factory(engine)
|
||||
try:
|
||||
with session_factory() as session:
|
||||
session.add(
|
||||
MemoryGraph(
|
||||
id="graph-local",
|
||||
organization_id="org-a",
|
||||
project_id="project-a",
|
||||
)
|
||||
)
|
||||
session.commit()
|
||||
|
||||
factory = make_local_entity_reader_factory(session_factory, organization_id="org-a")
|
||||
reader = factory("graph-local")
|
||||
assert isinstance(reader, LocalEntityReader)
|
||||
assert reader.repository.graph_id == "graph-local"
|
||||
assert reader.repository.organization_id == "org-a"
|
||||
assert reader.owns_session is True
|
||||
reader.close()
|
||||
finally:
|
||||
engine.dispose()
|
||||
62
backend/tests/test_local_report_tools.py
Normal file
62
backend/tests/test_local_report_tools.py
Normal file
@@ -0,0 +1,62 @@
|
||||
from types import SimpleNamespace
|
||||
|
||||
from sqlalchemy import create_engine
|
||||
|
||||
from app.db import Base, create_session_factory
|
||||
from app.models.memory import MemoryEdge, MemoryGraph, MemoryNode
|
||||
from app.services.memory_repository import SqlAlchemyMemoryRepository
|
||||
from app.services.memory_tools import LocalMemoryTools
|
||||
from app.services.report_agent import ReportAgent
|
||||
|
||||
|
||||
def test_report_agent_tools_use_local_memory_contract_without_zep():
|
||||
engine = create_engine("sqlite+pysqlite:///:memory:")
|
||||
Base.metadata.create_all(engine)
|
||||
factory = create_session_factory(engine)
|
||||
with factory() as session:
|
||||
session.add(MemoryGraph(id="graph-report", organization_id="org-a", project_id="project-a"))
|
||||
session.flush()
|
||||
repo = SqlAlchemyMemoryRepository(session, organization_id="org-a", graph_id="graph-report")
|
||||
alice = repo.upsert_node(
|
||||
canonical_name="Alice",
|
||||
labels=["Entity", "Person"],
|
||||
summary="Founder of Orbit",
|
||||
)
|
||||
orbit = repo.upsert_node(
|
||||
canonical_name="Orbit",
|
||||
labels=["Entity", "Project"],
|
||||
summary="A local project",
|
||||
)
|
||||
repo.upsert_edge(
|
||||
source_node_id=alice.id,
|
||||
target_node_id=orbit.id,
|
||||
relation="FOUNDED",
|
||||
fact="Alice founded Orbit",
|
||||
confidence=0.9,
|
||||
)
|
||||
session.commit()
|
||||
|
||||
tools = LocalMemoryTools(repo)
|
||||
agent = ReportAgent(
|
||||
graph_id="graph-report",
|
||||
simulation_id="simulation-a",
|
||||
simulation_requirement="Understand the project",
|
||||
llm_client=SimpleNamespace(),
|
||||
zep_tools=tools,
|
||||
)
|
||||
|
||||
quick = agent._execute_tool("quick_search", {"query": "Alice"})
|
||||
insight = agent._execute_tool("insight_forge", {"query": "Who founded Orbit?"})
|
||||
panorama = agent._execute_tool("panorama_search", {"query": "Orbit"})
|
||||
stats = agent._execute_tool("get_graph_statistics", {})
|
||||
summary = agent._execute_tool("get_entity_summary", {"entity_name": "Alice"})
|
||||
by_type = agent._execute_tool("get_entities_by_type", {"entity_type": "Person"})
|
||||
|
||||
assert "Alice founded Orbit" in quick
|
||||
assert "Alice founded Orbit" in insight
|
||||
assert "Orbit" in panorama
|
||||
assert '"node_count": 2' in stats
|
||||
assert "Alice" in summary
|
||||
assert "Alice" in by_type
|
||||
|
||||
engine.dispose()
|
||||
168
backend/tests/test_local_service_import.py
Normal file
168
backend/tests/test_local_service_import.py
Normal file
@@ -0,0 +1,168 @@
|
||||
import json
|
||||
import os
|
||||
import subprocess
|
||||
import sys
|
||||
|
||||
_ROOT = os.path.abspath(os.path.join(os.path.dirname(__file__), ".."))
|
||||
|
||||
|
||||
def _py_env():
|
||||
env = os.environ.copy()
|
||||
env["PYTHONPATH"] = _ROOT
|
||||
return env
|
||||
|
||||
|
||||
def _run_import_without_zep(module_name, symbol):
|
||||
script = f"""
|
||||
import builtins
|
||||
original_import = builtins.__import__
|
||||
|
||||
def guarded_import(name, *args, **kwargs):
|
||||
if name == 'zep_cloud' or name.startswith('zep_cloud.'):
|
||||
raise RuntimeError('zep_imported_during_local_service_import')
|
||||
return original_import(name, *args, **kwargs)
|
||||
|
||||
builtins.__import__ = guarded_import
|
||||
from app.services.{module_name} import {symbol} as imported_symbol
|
||||
print(imported_symbol.__name__)
|
||||
"""
|
||||
return subprocess.run(
|
||||
[sys.executable, "-c", script],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_py_env(),
|
||||
check=False,
|
||||
)
|
||||
|
||||
|
||||
def _run_import_explicit_zep(module_name, symbol):
|
||||
"""Fresh import that REQUIRES zep_cloud, to prove explicit Zep paths still resolve."""
|
||||
script = f"""
|
||||
import builtins
|
||||
original_import = builtins.__import__
|
||||
_zep_seen = []
|
||||
|
||||
def tracking_import(name, *args, **kwargs):
|
||||
if name == 'zep_cloud' or name.startswith('zep_cloud.'):
|
||||
_zep_seen.append(name)
|
||||
return original_import(name, *args, **kwargs)
|
||||
|
||||
builtins.__import__ = tracking_import
|
||||
from app.services.{module_name} import {symbol} as imported_symbol
|
||||
print(imported_symbol.__name__)
|
||||
print('ZEP_LOADED' if _zep_seen else 'NO_ZEP')
|
||||
"""
|
||||
return subprocess.run(
|
||||
[sys.executable, "-c", script],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_py_env(),
|
||||
check=False,
|
||||
)
|
||||
|
||||
|
||||
def test_local_service_import_does_not_eagerly_import_zep_client():
|
||||
result = _run_import_without_zep("local_graph_memory_updater", "LocalGraphMemoryUpdater")
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "LocalGraphMemoryUpdater"
|
||||
|
||||
|
||||
def test_local_simulation_manager_import_does_not_eagerly_import_zep_client():
|
||||
result = _run_import_without_zep("simulation_manager", "SimulationManager")
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "SimulationManager"
|
||||
|
||||
|
||||
def test_local_simulation_runner_import_does_not_eagerly_import_zep_client():
|
||||
result = _run_import_without_zep("simulation_runner", "SimulationRunner")
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "SimulationRunner"
|
||||
|
||||
|
||||
def test_local_report_agent_import_does_not_eagerly_import_zep_client():
|
||||
result = _run_import_without_zep("report_agent", "ReportAgent")
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "ReportAgent"
|
||||
|
||||
|
||||
def test_local_graph_builder_import_does_not_eagerly_import_zep_client():
|
||||
result = _run_import_without_zep("graph_builder", "GraphBuilderService")
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "GraphBuilderService"
|
||||
|
||||
|
||||
def test_explicit_zep_tools_service_import_still_resolves():
|
||||
# The lazy refactor must not break explicit Zep-only consumers.
|
||||
result = _run_import_explicit_zep("zep_tools", "ZepToolsService")
|
||||
assert result.returncode == 0, result.stderr
|
||||
lines = result.stdout.strip().splitlines()
|
||||
assert lines[0] == "ZepToolsService"
|
||||
assert "ZEP_LOADED" in lines
|
||||
|
||||
|
||||
def test_explicit_zep_entity_reader_import_still_resolves():
|
||||
result = _run_import_explicit_zep("zep_entity_reader", "ZepEntityReader")
|
||||
assert result.returncode == 0, result.stderr
|
||||
lines = result.stdout.strip().splitlines()
|
||||
assert lines[0] == "ZepEntityReader"
|
||||
assert "ZEP_LOADED" in lines
|
||||
|
||||
|
||||
_IDENTITY_PROBE = """
|
||||
import json
|
||||
from app.services.memory_activity import AgentActivity as Shared
|
||||
from app.services.zep_graph_memory_updater import ZepGraphMemoryUpdater
|
||||
|
||||
zep_alias = getattr(ZepGraphMemoryUpdater, "AgentActivity", None)
|
||||
if zep_alias is None:
|
||||
import app.services.zep_graph_memory_updater as zup
|
||||
zep_alias = getattr(zup, "AgentActivity", None)
|
||||
|
||||
print(json.dumps(
|
||||
{
|
||||
"zep_alias_resolved": zep_alias is not None,
|
||||
"shared_identity": zep_alias is Shared,
|
||||
"module": Shared.__module__,
|
||||
},
|
||||
ensure_ascii=False,
|
||||
))
|
||||
"""
|
||||
|
||||
|
||||
def test_local_and_zep_activity_contract_share_identity():
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", _IDENTITY_PROBE],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_py_env(),
|
||||
check=False,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
payload = json.loads(result.stdout.strip())
|
||||
# When the Zep updater aliases the shared class, identity must be exact.
|
||||
if payload["zep_alias_resolved"]:
|
||||
assert payload["shared_identity"] is True
|
||||
assert payload["module"] == "app.services.memory_activity"
|
||||
|
||||
|
||||
_GRAPH_BUILDER_SEAM_PROBE = """
|
||||
from app.services import graph_builder
|
||||
assert getattr(graph_builder, "Zep", None) is None
|
||||
from app.services import GraphBuilderService as Cls
|
||||
# The seam is a module-level sentinel that callers monkeypatch; it must remain
|
||||
# reassignable without constructing a real Zep client at import time.
|
||||
graph_builder.Zep = "sentinel"
|
||||
print(getattr(graph_builder, "Zep", None))
|
||||
"""
|
||||
|
||||
|
||||
def test_graph_builder_zep_seam_is_preserved_and_reassignable():
|
||||
result = subprocess.run(
|
||||
[sys.executable, "-c", _GRAPH_BUILDER_SEAM_PROBE],
|
||||
capture_output=True,
|
||||
text=True,
|
||||
env=_py_env(),
|
||||
check=False,
|
||||
)
|
||||
assert result.returncode == 0, result.stderr
|
||||
assert result.stdout.strip() == "sentinel"
|
||||
34
backend/tests/test_local_zep_boundary.py
Normal file
34
backend/tests/test_local_zep_boundary.py
Normal file
@@ -0,0 +1,34 @@
|
||||
import builtins
|
||||
|
||||
from app.services.local_graph_memory_updater import LocalGraphMemoryUpdater
|
||||
|
||||
|
||||
def test_local_activity_dict_path_does_not_import_zep_client(monkeypatch):
|
||||
updater = object.__new__(LocalGraphMemoryUpdater)
|
||||
captured = []
|
||||
setattr(updater, "add_activity", captured.append)
|
||||
|
||||
original_import = builtins.__import__
|
||||
|
||||
def reject_zep_import(name, *args, **kwargs):
|
||||
if name == "zep_cloud" or name.startswith("zep_cloud."):
|
||||
raise AssertionError("local runtime imported the Zep client")
|
||||
return original_import(name, *args, **kwargs)
|
||||
|
||||
monkeypatch.setattr(builtins, "__import__", reject_zep_import)
|
||||
|
||||
updater.add_activity_from_dict(
|
||||
{
|
||||
"agent_id": 7,
|
||||
"agent_name": "Alice",
|
||||
"action_type": "CREATE_POST",
|
||||
"action_args": {"content": "Local memory is durable."},
|
||||
"round": 2,
|
||||
"timestamp": "2026-08-24T00:00:00+00:00",
|
||||
},
|
||||
"twitter",
|
||||
)
|
||||
|
||||
assert len(captured) == 1
|
||||
assert captured[0].__class__.__module__ == "app.services.memory_activity"
|
||||
assert captured[0].to_episode_text().startswith("Alice:")
|
||||
Some files were not shown because too many files have changed in this diff Show More
Reference in New Issue
Block a user