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:
Kunthawat Greethong
2026-08-31 13:05:21 +07:00
parent 89d04e795b
commit 8b84378fe1
165 changed files with 15884 additions and 4001 deletions

View File

@@ -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

View 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 เดิมหาย.

View File

@@ -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
View 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

View File

@@ -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
View 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()})

View File

@@ -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
View 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}})

View File

@@ -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

View File

@@ -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)

View File

@@ -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)

View File

@@ -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:

View File

@@ -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
View 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)

View File

@@ -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',
]

View 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"
)

View 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
)

View 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()
)

View 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)

View File

@@ -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

View 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
View 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")

View 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()
)

View File

@@ -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]

View 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()
)

View 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",
]

View 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

View 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")

View 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")

View File

@@ -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)

View 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)

View 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()
)

View File

@@ -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)

View 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

View 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

View 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()

View 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)

View 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)

View 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}操作"

View 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

View 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",
]

View 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)

View 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,
)

View 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,
}

View File

@@ -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)

View 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

View 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")

View 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

View File

@@ -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,

View 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,
}

View File

@@ -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')

View File

@@ -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

View File

@@ -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()

View 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)

View File

@@ -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:

View File

@@ -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: 生成采访摘要

View 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)

View 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

View File

@@ -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
View 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()

View 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"}

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View 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")

View File

@@ -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]

View 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"]
}
}

View 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
}
]
}

View 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()

View 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()

View 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)

View 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

View 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()

View 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()

View 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()

View 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()

View 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()

View 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()

View 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)

View 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()

View 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()

View 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()

View 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()

View 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()

View 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()

View 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()

View 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()

View 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")

View 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()

View 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()

View 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()

View 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"

View 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