feat(ai): add tenant-safe hierarchical expense learning

This commit is contained in:
caoxiaozhu
2026-07-16 14:30:41 +08:00
parent 6bdf65bc24
commit ee88a36baf
65 changed files with 6909 additions and 232 deletions

View File

@@ -15,6 +15,7 @@ from sqlalchemy import (
String,
UniqueConstraint,
func,
text,
)
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.types import JSON
@@ -34,7 +35,7 @@ def _candidate_expires_at() -> datetime:
class MemoryEntry(Base):
"""受证据约束、可撤销的个人费用申请记忆。"""
"""受证据约束、可审计且可撤销的分层费用申请记忆。"""
__tablename__ = "memory_entries"
__table_args__ = (
@@ -59,9 +60,33 @@ class MemoryEntry(Base):
name="fk_memory_entries_tenant_superseded_by",
),
CheckConstraint(
"scope_type = 'user'",
"scope_type IN ('user', 'department', 'enterprise')",
name="ck_memory_entries_scope_type",
),
CheckConstraint(
"origin_type IN ('learned', 'admin_managed')",
name="ck_memory_entries_origin_type",
),
CheckConstraint(
"(scope_type = 'user' AND origin_type = 'learned') OR "
"(scope_type IN ('department', 'enterprise') "
"AND origin_type = 'admin_managed')",
name="ck_memory_entries_scope_origin",
),
CheckConstraint(
"scope_type != 'enterprise' OR scope_id = tenant_id",
name="ck_memory_entries_enterprise_scope",
),
CheckConstraint(
"(origin_type = 'learned' AND managed_by IS NULL "
"AND managed_at IS NULL AND management_reason IS NULL) OR "
"(origin_type = 'admin_managed' AND managed_by IS NOT NULL "
"AND length(trim(managed_by)) > 0 AND managed_at IS NOT NULL "
"AND management_reason IS NOT NULL "
"AND length(trim(management_reason)) > 0 "
"AND policy_version IS NOT NULL AND length(trim(policy_version)) > 0)",
name="ck_memory_entries_management_audit",
),
CheckConstraint(
"scene = 'travel_application'",
name="ck_memory_entries_scene",
@@ -120,6 +145,17 @@ class MemoryEntry(Base):
"superseded_by_id IS NULL OR superseded_by_id != id",
name="ck_memory_entries_not_self_superseded",
),
CheckConstraint(
"(management_request_id IS NULL AND management_payload_fingerprint IS NULL) OR "
"(management_request_id IS NOT NULL "
"AND management_payload_fingerprint IS NOT NULL)",
name="ck_memory_entries_management_idempotency_pair",
),
CheckConstraint(
"(revoke_request_id IS NULL AND revoke_payload_fingerprint IS NULL) OR "
"(revoke_request_id IS NOT NULL AND revoke_payload_fingerprint IS NOT NULL)",
name="ck_memory_entries_revoke_idempotency_pair",
),
Index(
"ix_memory_entries_scope_lookup",
"tenant_id",
@@ -136,6 +172,31 @@ class MemoryEntry(Base):
"candidate_expires_at",
"active_expires_at",
),
Index(
"uq_memory_entries_management_request",
"tenant_id",
"management_request_id",
unique=True,
),
Index(
"uq_memory_entries_revoke_request",
"tenant_id",
"revoke_request_id",
unique=True,
),
Index(
"uq_memory_entries_active_scope",
"tenant_id",
"scope_type",
"scope_id",
"scene",
"field_key",
unique=True,
postgresql_where=text(
"status = 'active' "
"AND scope_type IN ('department', 'enterprise')"
),
).ddl_if(dialect="postgresql"),
)
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=_new_id)
@@ -147,6 +208,16 @@ class MemoryEntry(Base):
server_default="user",
)
scope_id: Mapped[str] = mapped_column(String(120), nullable=False)
origin_type: Mapped[str] = mapped_column(
String(24),
nullable=False,
default="learned",
server_default="learned",
)
managed_by: Mapped[str | None] = mapped_column(String(255), nullable=True)
managed_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
management_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
policy_version: Mapped[str | None] = mapped_column(String(64), nullable=True)
scene: Mapped[str] = mapped_column(
String(50),
nullable=False,
@@ -212,6 +283,16 @@ class MemoryEntry(Base):
revoked_at: Mapped[datetime | None] = mapped_column(DateTime(timezone=True), nullable=True)
revoked_reason: Mapped[str | None] = mapped_column(String(255), nullable=True)
superseded_by_id: Mapped[str | None] = mapped_column(String(36), nullable=True)
management_request_id: Mapped[str | None] = mapped_column(String(120), nullable=True)
management_payload_fingerprint: Mapped[str | None] = mapped_column(
String(80),
nullable=True,
)
revoke_request_id: Mapped[str | None] = mapped_column(String(120), nullable=True)
revoke_payload_fingerprint: Mapped[str | None] = mapped_column(
String(80),
nullable=True,
)
created_at: Mapped[datetime] = mapped_column(
DateTime(timezone=True),
nullable=False,

View File

@@ -4,7 +4,7 @@ import uuid
from datetime import datetime
from typing import Any
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, func
from sqlalchemy import DateTime, ForeignKey, Index, String, Text, UniqueConstraint, func
from sqlalchemy.orm import Mapped, mapped_column
from sqlalchemy.types import JSON
@@ -20,12 +20,26 @@ class FewShotSample(Base):
__tablename__ = "few_shot_samples"
__table_args__ = (
UniqueConstraint(
"tenant_id",
"sample_key",
name="uq_few_shot_samples_tenant_key",
),
Index(
"ix_few_shot_samples_tenant_rule_lookup",
"tenant_id",
"scene",
"policy_ref",
"rule_version",
"status",
),
Index("ix_few_shot_samples_scene_label", "scene", "label"),
Index("ix_few_shot_samples_domain_risk_type", "domain", "risk_type"),
)
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
sample_key: Mapped[str] = mapped_column(String(160), unique=True, index=True)
tenant_id: Mapped[str] = mapped_column(String(64), default="default", index=True)
sample_key: Mapped[str] = mapped_column(String(160), index=True)
source_observation_id: Mapped[str | None] = mapped_column(
ForeignKey("risk_observations.id"),
nullable=True,
@@ -33,6 +47,8 @@ class FewShotSample(Base):
)
scene: Mapped[str] = mapped_column(String(50), default="risk_rule_generation", index=True)
policy_ref: Mapped[str] = mapped_column(String(160), default="", index=True)
rule_version: Mapped[str] = mapped_column(String(80), default="", index=True)
domain: Mapped[str] = mapped_column(String(50), default="", index=True)
risk_type: Mapped[str] = mapped_column(String(80), default="", index=True)
risk_level: Mapped[str] = mapped_column(String(20), default="")
@@ -45,7 +61,9 @@ class FewShotSample(Base):
vector_id: Mapped[str | None] = mapped_column(String(100), nullable=True)
status: Mapped[str] = mapped_column(String(20), default="active", index=True)
created_at: Mapped[datetime] = mapped_column(DateTime, default=func.now(), server_default=func.now())
created_at: Mapped[datetime] = mapped_column(
DateTime, default=func.now(), server_default=func.now()
)
updated_at: Mapped[datetime] = mapped_column(
DateTime,
default=func.now(),

View File

@@ -4,7 +4,17 @@ import uuid
from datetime import datetime
from typing import Any
from sqlalchemy import DateTime, Float, ForeignKey, Index, Integer, String, Text, func
from sqlalchemy import (
DateTime,
Float,
ForeignKey,
Index,
Integer,
String,
Text,
UniqueConstraint,
func,
)
from sqlalchemy.orm import Mapped, mapped_column, relationship
from sqlalchemy.types import JSON
@@ -14,18 +24,25 @@ from app.db.base_class import Base
class RiskObservation(Base):
__tablename__ = "risk_observations"
__table_args__ = (
UniqueConstraint(
"tenant_id",
"observation_key",
name="uq_risk_observations_tenant_key",
),
Index("ix_risk_observations_tenant_status", "tenant_id", "status", "created_at"),
Index("ix_risk_observations_subject", "subject_type", "subject_key"),
Index("ix_risk_observations_signal_level", "risk_signal", "risk_level"),
Index("ix_risk_observations_status_created", "status", "created_at"),
)
id: Mapped[str] = mapped_column(String(36), primary_key=True, default=lambda: str(uuid.uuid4()))
observation_key: Mapped[str] = mapped_column(String(160), unique=True, index=True)
tenant_id: Mapped[str] = mapped_column(String(64), default="default", index=True)
observation_key: Mapped[str] = mapped_column(String(160), index=True)
subject_type: Mapped[str] = mapped_column(String(50), index=True)
subject_key: Mapped[str] = mapped_column(String(160), index=True)
subject_label: Mapped[str] = mapped_column(String(160), default="")
claim_id: Mapped[str | None] = mapped_column(
ForeignKey("expense_claims.id"),
String(36),
nullable=True,
index=True,
)
@@ -66,7 +83,12 @@ class RiskObservation(Base):
onupdate=func.now(),
)
claim = relationship("ExpenseClaim", foreign_keys=[claim_id])
claim = relationship(
"ExpenseClaim",
primaryjoin="foreign(RiskObservation.claim_id) == ExpenseClaim.id",
foreign_keys=[claim_id],
viewonly=True,
)
feedback_items = relationship(
"RiskObservationFeedback",
back_populates="observation",