Files
company-ai-platform/app/modules/ai_memory/service.py
JiuContinent db751f03b4 ```
refactor(Dockerfile): 使用requirements.txt替代硬编码依赖

将Dockerfile中的硬编码pip包列表替换为通过requirements.txt文件安装,
提高依赖管理的灵活性和可维护性。

feat(scheduling): 移除内置APScheduler,采用独立调度系统

移除app/core/background/scheduler.py中原来的APScheduler实现,
改为使用新的应用级调度系统app.application.scheduling。

refactor(task_queue): 调整任务队列模块结构和导入路径

将任务队列相关常量从app.core.background.task_queue.constants迁移至
app.tasks.constants,并更新所有相关导入路径和引用。

refactor(events): 将事件服务重构为独立的应用层组件

将事件分发逻辑从核心层迁移到应用层,使用app.application.events.EventDispatchService
替代原有的app.modules.events.services.EventService。

feat(ai_memory): 增强AI记忆自动写入的安全策略

新增ai_memory_blocked_content_terms配置项用于阻止敏感内容,
添加TTL过期机制控制自动写入条目的生命周期。

fix(security): 强化生产环境安全验证机制

增加model_validator确保生产环境中数据库连接、API密钥、CORS设置等
关键安全配置符合要求。

feat(risks): 优化风险事件操作动作的外键约束

为RiskEventAction模型的风险事件ID字段添加外键约束,
防止孤立记录并增强数据完整性。

refactor(audit): 优化审计服务方法命名和事务处理

将AuditService的log方法重命名为record以反映其阶段行为,
并调整事务提交时机以提高性能。

feat(events): 增强领域事件并发处理和响应模型

添加事件锁定机制防止重复处理,更新API响应模型以提供
更准确的数据类型定义。
```
2026-07-15 16:36:42 +08:00

449 lines
16 KiB
Python

from datetime import datetime, timedelta
from typing import Any
from uuid import uuid4
from fastapi import HTTPException, status
from sqlalchemy import func, or_, select
from sqlalchemy.orm import Session
from app.core.config import get_settings
from app.core.constants import ActorValue
from app.core.http.pagination import bounded_limit
from app.core.utils.time import utc_now
from app.modules.ai_memory.constants import (
AI_MEMORY_CODE_PREFIX,
AI_MEMORY_MAX_CONTENT_LENGTH,
AI_MEMORY_MAX_SUMMARY_LENGTH,
AI_MEMORY_MIN_AUTO_WRITE_LENGTH,
AIMemoryPayloadKey,
AIMemoryScope,
AIMemorySource,
AIMemoryStatus,
AIMemoryText,
AI_USER_RULE_MAX_PRIORITY,
AI_USER_RULE_MIN_PRIORITY,
)
from app.modules.ai_memory.models import AIMemoryEntry
from app.modules.audit.constants import (
AuditAction,
AuditRiskLevel,
AuditSource,
AuditTargetType,
)
from app.modules.audit.schemas import AuditLogCreate
from app.modules.audit.service import AuditService
from app.modules.business.service import serialize_model
from app.modules.events.constants import EventAggregateType, EventSource, EventType
from app.modules.events.services import EventService
class AIMemoryService:
"""Store and recall audited local AI memory for read-only operations."""
def __init__(self, db: Session):
self.db = db
self.audit = AuditService(db)
def list_entries(
self,
scope: str | None = None,
subject: str | None = None,
status_filter: str = AIMemoryStatus.ACTIVE,
limit: int = 100,
) -> list[dict[str, Any]]:
stmt = (
select(AIMemoryEntry)
.where(AIMemoryEntry.status == status_filter)
.order_by(AIMemoryEntry.importance.desc(), AIMemoryEntry.id.desc())
.limit(bounded_limit(limit))
)
if scope:
stmt = stmt.where(AIMemoryEntry.scope == scope)
if subject:
stmt = stmt.where(AIMemoryEntry.subject == subject)
return [serialize_model(item) for item in self.db.execute(stmt).scalars()]
def recall(
self,
query: str,
scope: str = AIMemoryScope.GLOBAL,
subject: str | None = None,
limit: int | None = None,
actor: str = ActorValue.API,
) -> list[dict[str, Any]]:
settings = get_settings()
if not settings.ai_memory_enabled:
return []
limit_value = bounded_limit(limit or settings.ai_memory_recall_limit)
now = utc_now()
stmt = (
select(AIMemoryEntry)
.where(
AIMemoryEntry.status == AIMemoryStatus.ACTIVE,
or_(AIMemoryEntry.expires_at.is_(None), AIMemoryEntry.expires_at > now),
AIMemoryEntry.scope.in_({AIMemoryScope.GLOBAL, scope}),
)
.order_by(AIMemoryEntry.importance.desc(), AIMemoryEntry.id.desc())
.limit(limit_value * 3)
)
if subject:
stmt = stmt.where(
or_(
AIMemoryEntry.subject == subject,
AIMemoryEntry.scope == AIMemoryScope.GLOBAL,
)
)
candidates = list(self.db.execute(stmt).scalars())
items = [item for item in candidates if _matches_query(item, query)]
if not items:
items = candidates[:limit_value]
items = items[:limit_value]
for item in items:
item.last_used_at = now
result = [serialize_model(item) for item in items]
self.audit.record(
AuditLogCreate(
actor=actor,
source=AuditSource.AI_MEMORY,
action=AuditAction.AI_MEMORY_RECALL,
target_type=AuditTargetType.AI_MEMORY,
risk_level=AuditRiskLevel.LOW,
request_payload={
AIMemoryPayloadKey.QUERY: query,
AIMemoryPayloadKey.SCOPE: scope,
AIMemoryPayloadKey.SUBJECT: subject,
AIMemoryPayloadKey.LIMIT: limit_value,
},
response_payload={AIMemoryPayloadKey.COUNT: len(result)},
)
)
self.db.commit()
return result
def auto_write(
self,
prompt: str,
context: dict[str, Any],
answer: str,
actor: str = ActorValue.API,
) -> AIMemoryEntry | None:
settings = get_settings()
if not settings.ai_memory_enabled or not settings.ai_memory_auto_write_enabled:
return None
content = _build_memory_content(prompt, context, answer)
if len(content) < AI_MEMORY_MIN_AUTO_WRITE_LENGTH:
return None
scope = str(context.get(AIMemoryPayloadKey.SCOPE) or AIMemoryText.DEFAULT_SCOPE)
subject = str(context.get(AIMemoryPayloadKey.SUBJECT) or AIMemoryText.DEFAULT_SUBJECT)
if _contains_forbidden_value(
{
"prompt": prompt,
"context": context,
"answer": answer,
},
settings.ai_memory_forbidden_keys,
):
record = self._create_entry(
scope=scope,
subject=subject,
content=str(AIMemoryText.REJECTED_SECRET),
summary=str(AIMemoryText.REJECTED_SECRET),
tags=[str(AIMemoryText.AUTO_TAG)],
source=AIMemorySource.AUTO,
importance=0,
status_value=AIMemoryStatus.REJECTED,
actor=actor,
expires_at=utc_now()
+ timedelta(days=settings.ai_memory_auto_write_ttl_days),
)
return record
if _contains_blocked_content(content, settings.ai_memory_blocked_content_terms):
return self._create_entry(
scope=scope,
subject=subject,
content=str(AIMemoryText.REJECTED_SENSITIVE_FACT),
summary=str(AIMemoryText.REJECTED_SENSITIVE_FACT),
tags=[str(AIMemoryText.AUTO_TAG)],
source=AIMemorySource.AUTO,
importance=0,
status_value=AIMemoryStatus.REJECTED,
actor=actor,
expires_at=utc_now()
+ timedelta(days=settings.ai_memory_auto_write_ttl_days),
)
summary = _truncate(answer, AI_MEMORY_MAX_SUMMARY_LENGTH)
record = self._create_entry(
scope=scope,
subject=subject,
content=_truncate(content, AI_MEMORY_MAX_CONTENT_LENGTH),
summary=summary,
tags=[str(AIMemoryText.AUTO_TAG)],
source=AIMemorySource.AUTO,
importance=1,
status_value=AIMemoryStatus.ACTIVE,
actor=actor,
expires_at=utc_now() + timedelta(days=settings.ai_memory_auto_write_ttl_days),
)
return record
def list_rules(
self,
scope: str | None = None,
subject: str | None = None,
status_filter: str | None = None,
limit: int = 100,
) -> list[dict[str, Any]]:
stmt = (
select(AIMemoryEntry)
.where(AIMemoryEntry.source == AIMemorySource.USER_RULE)
.order_by(AIMemoryEntry.importance.desc(), AIMemoryEntry.id.asc())
.limit(bounded_limit(limit))
)
if scope:
stmt = stmt.where(AIMemoryEntry.scope == scope)
if subject:
stmt = stmt.where(AIMemoryEntry.subject == subject)
if status_filter:
stmt = stmt.where(AIMemoryEntry.status == status_filter)
return [serialize_model(item) for item in self.db.execute(stmt).scalars()]
def active_rules(
self,
scope: str = AIMemoryScope.GLOBAL,
subject: str | None = None,
limit: int = 50,
) -> list[dict[str, Any]]:
stmt = (
select(AIMemoryEntry)
.where(
AIMemoryEntry.source == AIMemorySource.USER_RULE,
AIMemoryEntry.status == AIMemoryStatus.ACTIVE,
AIMemoryEntry.scope.in_({AIMemoryScope.GLOBAL, scope}),
)
.order_by(AIMemoryEntry.importance.desc(), AIMemoryEntry.id.asc())
.limit(min(bounded_limit(limit), 50))
)
if subject:
stmt = stmt.where(
or_(
AIMemoryEntry.scope == AIMemoryScope.GLOBAL,
AIMemoryEntry.subject == subject,
)
)
return [
{
"code": item.code,
"scope": item.scope,
"subject": item.subject,
"rule": item.content,
"priority": item.importance,
}
for item in self.db.execute(stmt).scalars()
]
def create_rule(
self,
content: str,
scope: str,
subject: str,
priority: int,
tags: list[str] | None,
actor: str,
) -> dict[str, Any]:
self._validate_rule(content, priority)
record = self._create_entry(
scope=scope,
subject=subject,
content=_truncate(content.strip(), AI_MEMORY_MAX_CONTENT_LENGTH),
summary=_truncate(content.strip(), AI_MEMORY_MAX_SUMMARY_LENGTH),
tags=["user-rule", *(tags or [])],
source=AIMemorySource.USER_RULE,
importance=priority,
status_value=AIMemoryStatus.ACTIVE,
actor=actor,
audit_action=AuditAction.AI_RULE_CREATE,
)
return serialize_model(record)
def update_rule(
self,
code: str,
content: str | None,
priority: int | None,
tags: list[str] | None,
enabled: bool | None,
actor: str,
) -> dict[str, Any]:
record = self.db.execute(
select(AIMemoryEntry).where(
AIMemoryEntry.code == code,
AIMemoryEntry.source == AIMemorySource.USER_RULE,
)
).scalar_one_or_none()
if record is None:
raise HTTPException(status_code=status.HTTP_404_NOT_FOUND, detail="AI rule not found")
if content is not None:
self._validate_rule(content, priority or record.importance)
record.content = _truncate(content.strip(), AI_MEMORY_MAX_CONTENT_LENGTH)
record.summary = _truncate(content.strip(), AI_MEMORY_MAX_SUMMARY_LENGTH)
if priority is not None:
self._validate_rule(record.content, priority)
record.importance = priority
if tags is not None:
record.tags = ["user-rule", *tags]
if enabled is not None:
record.status = AIMemoryStatus.ACTIVE if enabled else AIMemoryStatus.ARCHIVED
self.audit.record(
AuditLogCreate(
actor=actor,
source=AuditSource.AI_MEMORY,
action=AuditAction.AI_RULE_UPDATE,
target_type=AuditTargetType.AI_MEMORY,
target_id=record.code,
risk_level=AuditRiskLevel.MEDIUM,
response_payload={
AIMemoryPayloadKey.CODE: record.code,
AIMemoryPayloadKey.STATUS: record.status,
AIMemoryPayloadKey.IMPORTANCE: record.importance,
},
)
)
self.db.commit()
self.db.refresh(record)
return serialize_model(record)
def _validate_rule(self, content: str, priority: int) -> None:
if not AI_USER_RULE_MIN_PRIORITY <= priority <= AI_USER_RULE_MAX_PRIORITY:
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="AI rule priority is out of range",
)
if _contains_forbidden_value(content, get_settings().ai_memory_forbidden_keys):
raise HTTPException(
status_code=status.HTTP_422_UNPROCESSABLE_ENTITY,
detail="AI rule contains secret-like content",
)
def count_by_status(self) -> dict[str, int]:
rows = self.db.execute(
select(AIMemoryEntry.status, func.count()).group_by(AIMemoryEntry.status)
).all()
return {str(status_value): int(count) for status_value, count in rows}
def _create_entry(
self,
scope: str,
subject: str,
content: str,
summary: str | None,
tags: list[str],
source: str,
importance: int,
status_value: str,
actor: str,
audit_action: str = AuditAction.AI_MEMORY_WRITE,
expires_at: datetime | None = None,
) -> AIMemoryEntry:
record = AIMemoryEntry(
code=(
f"{AI_MEMORY_CODE_PREFIX}-{utc_now():%Y%m%d%H%M%S%f}-"
f"{uuid4().hex[:8]}"
),
scope=scope,
subject=subject,
content=content,
summary=summary,
tags=tags,
source=source,
importance=importance,
status=status_value,
actor=actor,
expires_at=expires_at,
)
self.db.add(record)
self.db.flush()
self.audit.record(
AuditLogCreate(
actor=actor,
source=AuditSource.AI_MEMORY,
action=audit_action,
target_type=AuditTargetType.AI_MEMORY,
target_id=record.code,
risk_level=AuditRiskLevel.LOW,
request_payload={
AIMemoryPayloadKey.SCOPE: scope,
AIMemoryPayloadKey.SUBJECT: subject,
AIMemoryPayloadKey.SOURCE: source,
AIMemoryPayloadKey.STATUS: status_value,
},
response_payload={AIMemoryPayloadKey.CODE: record.code},
)
)
EventService(self.db).enqueue(
event_type=EventType.AI_MEMORY_WRITTEN,
source=EventSource.AI_MEMORY,
aggregate_type=EventAggregateType.AI_MEMORY_ENTRY,
aggregate_id=record.code,
actor=actor,
payload={
AIMemoryPayloadKey.CODE: record.code,
AIMemoryPayloadKey.SCOPE: scope,
AIMemoryPayloadKey.SUBJECT: subject,
AIMemoryPayloadKey.STATUS: status_value,
},
idempotency_key=f"ai-memory:{record.code}",
)
self.db.commit()
self.db.refresh(record)
return record
def _build_memory_content(prompt: str, context: dict[str, Any], answer: str) -> str:
context_text = ", ".join(
f"{key}={value}" for key, value in sorted(context.items(), key=lambda item: str(item[0]))
)
return f"prompt: {prompt}\ncontext: {context_text}\nanswer: {answer}"
def _contains_forbidden_value(value: Any, forbidden_keys: list[str]) -> bool:
forbidden = {item.lower() for item in forbidden_keys}
if isinstance(value, dict):
for key, item in value.items():
if str(key).lower() in forbidden:
return True
if _contains_forbidden_value(item, forbidden_keys):
return True
return False
if isinstance(value, (list, tuple, set)):
return any(_contains_forbidden_value(item, forbidden_keys) for item in value)
if isinstance(value, str):
lowered = value.lower()
return any(item in lowered for item in forbidden)
return False
def _contains_blocked_content(value: str, blocked_terms: list[str]) -> bool:
lowered = value.lower()
return any(term.lower() in lowered for term in blocked_terms if term.strip())
def _matches_query(entry: AIMemoryEntry, query: str) -> bool:
query_text = query.lower().strip()
if not query_text:
return True
text = " ".join(
[
entry.subject or "",
entry.content or "",
entry.summary or "",
" ".join(str(item) for item in (entry.tags or [])),
]
).lower()
return any(token in text for token in query_text.split())
def _truncate(value: str, max_length: int) -> str:
if len(value) <= max_length:
return value
return value[:max_length]