memlord 0.2.3__py3-none-any.whl
This diff represents the content of publicly available package versions that have been released to one of the supported registries. The information contained in this diff is provided for informational purposes only and reflects changes between package versions as they appear in their respective public registries.
- memlord/__init__.py +0 -0
- memlord/auth.py +42 -0
- memlord/config.py +34 -0
- memlord/dao/__init__.py +4 -0
- memlord/dao/email_token.py +65 -0
- memlord/dao/memory.py +289 -0
- memlord/dao/user.py +105 -0
- memlord/dao/workspace.py +411 -0
- memlord/db.py +40 -0
- memlord/embeddings.py +64 -0
- memlord/main.py +79 -0
- memlord/models/__init__.py +10 -0
- memlord/models/base.py +17 -0
- memlord/models/email_token.py +22 -0
- memlord/models/memory.py +43 -0
- memlord/models/memory_tag.py +14 -0
- memlord/models/oauth_client.py +13 -0
- memlord/models/revoked_token.py +10 -0
- memlord/models/schema_version.py +10 -0
- memlord/models/tag.py +10 -0
- memlord/models/user.py +16 -0
- memlord/models/workspace.py +61 -0
- memlord/oauth.py +703 -0
- memlord/onnx/.gitignore +1 -0
- memlord/schemas/__init__.py +9 -0
- memlord/schemas/delete.py +6 -0
- memlord/schemas/list_memories.py +23 -0
- memlord/schemas/memory_type.py +8 -0
- memlord/schemas/recall.py +14 -0
- memlord/schemas/search.py +25 -0
- memlord/schemas/store.py +19 -0
- memlord/schemas/update.py +10 -0
- memlord/schemas/user.py +8 -0
- memlord/schemas/workspace.py +28 -0
- memlord/search.py +144 -0
- memlord/server.py +46 -0
- memlord/templates/_invite_link.html +5 -0
- memlord/templates/_memory_content.html +120 -0
- memlord/templates/base.html +775 -0
- memlord/templates/forgot_password.html +71 -0
- memlord/templates/icon.svg +12 -0
- memlord/templates/index.html +260 -0
- memlord/templates/login.html +76 -0
- memlord/templates/memory.html +13 -0
- memlord/templates/register.html +86 -0
- memlord/templates/reset_password.html +77 -0
- memlord/templates/search.html +48 -0
- memlord/templates/verify_email.html +62 -0
- memlord/templates/workspace_detail.html +125 -0
- memlord/templates/workspace_join.html +27 -0
- memlord/templates/workspace_new.html +35 -0
- memlord/templates/workspaces.html +39 -0
- memlord/tools/__init__.py +10 -0
- memlord/tools/delete.py +25 -0
- memlord/tools/get_memory.py +23 -0
- memlord/tools/list_memories.py +95 -0
- memlord/tools/move.py +37 -0
- memlord/tools/recall.py +103 -0
- memlord/tools/retrieve.py +78 -0
- memlord/tools/search_by_tag.py +87 -0
- memlord/tools/store.py +57 -0
- memlord/tools/update.py +43 -0
- memlord/tools/workspaces.py +19 -0
- memlord/ui/__init__.py +12 -0
- memlord/ui/base.py +344 -0
- memlord/ui/data.py +128 -0
- memlord/ui/login.py +241 -0
- memlord/ui/utils.py +86 -0
- memlord/ui/workspaces.py +241 -0
- memlord/utils/__init__.py +0 -0
- memlord/utils/dt.py +5 -0
- memlord/utils/inject_client_id.py +49 -0
- memlord/utils/mail_send.py +25 -0
- memlord-0.2.3.dist-info/METADATA +683 -0
- memlord-0.2.3.dist-info/RECORD +79 -0
- memlord-0.2.3.dist-info/WHEEL +4 -0
- memlord-0.2.3.dist-info/entry_points.txt +2 -0
- memlord-0.2.3.dist-info/licenses/LICENSE +661 -0
- memlord-0.2.3.dist-info/licenses/LICENSE-COMMERCIAL +35 -0
memlord/__init__.py
ADDED
|
File without changes
|
memlord/auth.py
ADDED
|
@@ -0,0 +1,42 @@
|
|
|
1
|
+
from contextlib import asynccontextmanager
|
|
2
|
+
|
|
3
|
+
import bcrypt
|
|
4
|
+
from fastmcp.server.dependencies import get_access_token
|
|
5
|
+
from sqlalchemy import select
|
|
6
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
7
|
+
|
|
8
|
+
from fastmcp.dependencies import Depends as MCPDepends
|
|
9
|
+
|
|
10
|
+
from memlord.config import settings
|
|
11
|
+
from memlord.db import MCPSessionDep
|
|
12
|
+
from memlord.models.oauth_client import OAuthClient
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def hash_password(password: str) -> str:
|
|
16
|
+
return bcrypt.hashpw(password.encode(), bcrypt.gensalt()).decode()
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
def verify_password(plain: str, hashed: str) -> bool:
|
|
20
|
+
return bcrypt.checkpw(plain.encode(), hashed.encode())
|
|
21
|
+
|
|
22
|
+
|
|
23
|
+
async def _current_user_gen(
|
|
24
|
+
s: AsyncSession = MCPSessionDep, # type: ignore[assignment]
|
|
25
|
+
):
|
|
26
|
+
access_token = get_access_token()
|
|
27
|
+
if access_token is None:
|
|
28
|
+
if settings.stdio_user_id is None:
|
|
29
|
+
raise PermissionError("Authentication required")
|
|
30
|
+
yield settings.stdio_user_id
|
|
31
|
+
return
|
|
32
|
+
user_id = await s.scalar(
|
|
33
|
+
select(OAuthClient.user_id).where(
|
|
34
|
+
OAuthClient.client_id == access_token.client_id
|
|
35
|
+
)
|
|
36
|
+
)
|
|
37
|
+
if user_id is None:
|
|
38
|
+
raise PermissionError("Unauthenticated")
|
|
39
|
+
yield user_id
|
|
40
|
+
|
|
41
|
+
|
|
42
|
+
MCPUserDep = MCPDepends(asynccontextmanager(_current_user_gen))
|
memlord/config.py
ADDED
|
@@ -0,0 +1,34 @@
|
|
|
1
|
+
from pathlib import Path
|
|
2
|
+
|
|
3
|
+
from pydantic import EmailStr, Field
|
|
4
|
+
from pydantic_settings import BaseSettings, SettingsConfigDict
|
|
5
|
+
|
|
6
|
+
|
|
7
|
+
class Settings(BaseSettings):
|
|
8
|
+
model_config = SettingsConfigDict(
|
|
9
|
+
env_prefix="MEMLORD_", env_file=".env", extra="ignore"
|
|
10
|
+
)
|
|
11
|
+
|
|
12
|
+
db_url: str = "postgresql+asyncpg://postgres:postgres@localhost/memlord"
|
|
13
|
+
db_echo: bool = False
|
|
14
|
+
|
|
15
|
+
model_dir: Path = Path("src/memlord/onnx")
|
|
16
|
+
host: str = "0.0.0.0"
|
|
17
|
+
port: int = 8000
|
|
18
|
+
base_url: str = "http://localhost:8000"
|
|
19
|
+
rrf_k: int = 60
|
|
20
|
+
default_limit: int = 10
|
|
21
|
+
sim_threshold: float = Field(0.25, ge=0.0, le=1.0)
|
|
22
|
+
dedup_threshold: float = Field(0.85, ge=0.0, le=1.0)
|
|
23
|
+
oauth_jwt_secret: str = "memlord-dev-secret-please-change"
|
|
24
|
+
stdio_user_id: int | None = Field(None, description="use for stdio mode")
|
|
25
|
+
|
|
26
|
+
smtp_host: str | None = None
|
|
27
|
+
smtp_port: int = 587
|
|
28
|
+
smtp_user: str | None = None
|
|
29
|
+
smtp_password: str | None = None
|
|
30
|
+
smtp_from: EmailStr | None = None
|
|
31
|
+
smtp_tls: bool = True
|
|
32
|
+
|
|
33
|
+
|
|
34
|
+
settings = Settings()
|
memlord/dao/__init__.py
ADDED
|
@@ -0,0 +1,65 @@
|
|
|
1
|
+
import hashlib
|
|
2
|
+
import secrets
|
|
3
|
+
from datetime import timedelta
|
|
4
|
+
|
|
5
|
+
from sqlalchemy import delete, insert, select
|
|
6
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
7
|
+
|
|
8
|
+
from memlord.models.email_token import EmailToken, TokenPurpose
|
|
9
|
+
from memlord.utils.dt import utcnow
|
|
10
|
+
|
|
11
|
+
_TTL: dict[TokenPurpose, timedelta] = {
|
|
12
|
+
TokenPurpose.verify: timedelta(hours=24),
|
|
13
|
+
TokenPurpose.reset: timedelta(hours=1),
|
|
14
|
+
}
|
|
15
|
+
|
|
16
|
+
|
|
17
|
+
def _hash(token: str) -> str:
|
|
18
|
+
return hashlib.sha256(token.encode()).hexdigest()
|
|
19
|
+
|
|
20
|
+
|
|
21
|
+
class EmailTokenDao:
|
|
22
|
+
def __init__(self, s: AsyncSession) -> None:
|
|
23
|
+
self._s = s
|
|
24
|
+
|
|
25
|
+
async def create(self, user_id: int, purpose: TokenPurpose) -> str:
|
|
26
|
+
"""Create a new token, delete any existing one for same user+purpose, return raw token."""
|
|
27
|
+
await self._s.execute(
|
|
28
|
+
delete(EmailToken).where(
|
|
29
|
+
EmailToken.user_id == user_id,
|
|
30
|
+
EmailToken.purpose == purpose,
|
|
31
|
+
)
|
|
32
|
+
)
|
|
33
|
+
raw = secrets.token_urlsafe(32)
|
|
34
|
+
await self._s.execute(
|
|
35
|
+
insert(EmailToken).values(
|
|
36
|
+
token_hash=_hash(raw),
|
|
37
|
+
user_id=user_id,
|
|
38
|
+
purpose=purpose,
|
|
39
|
+
expires_at=utcnow() + _TTL[purpose],
|
|
40
|
+
)
|
|
41
|
+
)
|
|
42
|
+
return raw
|
|
43
|
+
|
|
44
|
+
async def consume(self, raw: str, purpose: TokenPurpose) -> int | None:
|
|
45
|
+
"""Validate token; on success delete it and return user_id, else return None."""
|
|
46
|
+
row = (
|
|
47
|
+
(
|
|
48
|
+
await self._s.execute(
|
|
49
|
+
select(EmailToken.user_id, EmailToken.expires_at).where(
|
|
50
|
+
EmailToken.token_hash == _hash(raw),
|
|
51
|
+
EmailToken.purpose == purpose,
|
|
52
|
+
)
|
|
53
|
+
)
|
|
54
|
+
)
|
|
55
|
+
.mappings()
|
|
56
|
+
.one_or_none()
|
|
57
|
+
)
|
|
58
|
+
if row is None:
|
|
59
|
+
return None
|
|
60
|
+
if utcnow() > row["expires_at"]:
|
|
61
|
+
return None
|
|
62
|
+
await self._s.execute(
|
|
63
|
+
delete(EmailToken).where(EmailToken.token_hash == _hash(raw))
|
|
64
|
+
)
|
|
65
|
+
return row["user_id"]
|
memlord/dao/memory.py
ADDED
|
@@ -0,0 +1,289 @@
|
|
|
1
|
+
from pgvector.sqlalchemy import Vector
|
|
2
|
+
from sqlalchemy import bindparam, delete, Float, insert, select, update
|
|
3
|
+
from sqlalchemy.dialects.postgresql import insert as pg_insert
|
|
4
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
5
|
+
|
|
6
|
+
from memlord.config import settings
|
|
7
|
+
from memlord.dao.workspace import WorkspaceDao
|
|
8
|
+
from memlord.embeddings import embed
|
|
9
|
+
from memlord.models import Memory, MemoryTag, Tag
|
|
10
|
+
from memlord.schemas import MemoryListItem, MemoryType
|
|
11
|
+
|
|
12
|
+
_UNSET = object()
|
|
13
|
+
|
|
14
|
+
|
|
15
|
+
def _embed_text(content: str, tags: set[str]) -> str:
|
|
16
|
+
return f"{content} {' '.join(sorted(tags))}" if tags else content
|
|
17
|
+
|
|
18
|
+
|
|
19
|
+
class MemoryDao:
|
|
20
|
+
def __init__(self, s: AsyncSession, uid: int) -> None:
|
|
21
|
+
self._s = s
|
|
22
|
+
self._uid = uid
|
|
23
|
+
|
|
24
|
+
async def _upsert_tags(self, memory_id: int, tags: set[str]) -> None:
|
|
25
|
+
for tag_name in tags:
|
|
26
|
+
normalized = tag_name.lower().strip()
|
|
27
|
+
if not normalized:
|
|
28
|
+
continue
|
|
29
|
+
await self._s.execute(
|
|
30
|
+
pg_insert(Tag).values(name=normalized).on_conflict_do_nothing()
|
|
31
|
+
)
|
|
32
|
+
tag_id = await self._s.scalar(select(Tag.id).where(Tag.name == normalized))
|
|
33
|
+
await self._s.execute(
|
|
34
|
+
pg_insert(MemoryTag)
|
|
35
|
+
.values(memory_id=memory_id, tag_id=tag_id)
|
|
36
|
+
.on_conflict_do_nothing()
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
async def _fetch_tag_names(self, memory_id: int) -> list[str]:
|
|
40
|
+
rows = await self._s.execute(
|
|
41
|
+
select(Tag.name)
|
|
42
|
+
.join(MemoryTag, MemoryTag.tag_id == Tag.id)
|
|
43
|
+
.where(MemoryTag.memory_id == memory_id)
|
|
44
|
+
)
|
|
45
|
+
return [row[0] for row in rows.fetchall()]
|
|
46
|
+
|
|
47
|
+
async def _cleanup_orphan_tags(self) -> None:
|
|
48
|
+
await self._s.execute(delete(Tag).where(~Tag.id.in_(select(MemoryTag.tag_id))))
|
|
49
|
+
|
|
50
|
+
async def _replace_tags(self, memory_id: int, tags: set[str]) -> None:
|
|
51
|
+
await self._s.execute(delete(MemoryTag).where(MemoryTag.memory_id == memory_id))
|
|
52
|
+
await self._upsert_tags(memory_id, tags)
|
|
53
|
+
await self._cleanup_orphan_tags()
|
|
54
|
+
|
|
55
|
+
async def _check_near_duplicate(
|
|
56
|
+
self, vector: list[float], workspace_id: int
|
|
57
|
+
) -> None:
|
|
58
|
+
"""Raise ValueError if a near-duplicate exists in the workspace."""
|
|
59
|
+
vec_param = bindparam("vec", type_=Vector(384))
|
|
60
|
+
distance_expr = Memory.embedding.op("<=>", return_type=Float)(vec_param)
|
|
61
|
+
dup_row = (
|
|
62
|
+
(
|
|
63
|
+
await self._s.execute(
|
|
64
|
+
select(Memory.id, distance_expr.label("distance"))
|
|
65
|
+
.where(
|
|
66
|
+
Memory.embedding.isnot(None),
|
|
67
|
+
Memory.workspace_id == workspace_id,
|
|
68
|
+
)
|
|
69
|
+
.order_by(distance_expr)
|
|
70
|
+
.limit(1),
|
|
71
|
+
{"vec": vector},
|
|
72
|
+
)
|
|
73
|
+
)
|
|
74
|
+
.mappings()
|
|
75
|
+
.one_or_none()
|
|
76
|
+
)
|
|
77
|
+
if dup_row is None:
|
|
78
|
+
return
|
|
79
|
+
similarity = 1.0 - dup_row["distance"]
|
|
80
|
+
if similarity >= settings.dedup_threshold:
|
|
81
|
+
raise ValueError(
|
|
82
|
+
f"Near-duplicate found (id={dup_row['id']}, similarity={round(similarity, 4):.4f}). "
|
|
83
|
+
f"Review with get_memory({dup_row['id']}). Pass force=True to store anyway."
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
async def _personal_workspace_id(self) -> int:
|
|
87
|
+
ws = await WorkspaceDao(self._s).get_personal(self._uid)
|
|
88
|
+
return ws.id
|
|
89
|
+
|
|
90
|
+
async def _accessible_workspace_ids(self) -> list[int]:
|
|
91
|
+
return await WorkspaceDao(self._s).get_accessible_workspace_ids(self._uid)
|
|
92
|
+
|
|
93
|
+
async def create(
|
|
94
|
+
self,
|
|
95
|
+
content: str,
|
|
96
|
+
memory_type: MemoryType,
|
|
97
|
+
metadata: dict,
|
|
98
|
+
tags: set[str],
|
|
99
|
+
workspace_id: int | None = None,
|
|
100
|
+
force: bool = False,
|
|
101
|
+
) -> tuple[int, bool]:
|
|
102
|
+
if workspace_id is None:
|
|
103
|
+
workspace_id = await self._personal_workspace_id()
|
|
104
|
+
|
|
105
|
+
memory_id = await self._s.scalar(
|
|
106
|
+
select(Memory.id).where(
|
|
107
|
+
Memory.content == content,
|
|
108
|
+
Memory.workspace_id == workspace_id,
|
|
109
|
+
)
|
|
110
|
+
)
|
|
111
|
+
if memory_id is not None:
|
|
112
|
+
return memory_id, False
|
|
113
|
+
|
|
114
|
+
vector = await embed(_embed_text(content, tags or []))
|
|
115
|
+
|
|
116
|
+
if not force:
|
|
117
|
+
await self._check_near_duplicate(vector, workspace_id)
|
|
118
|
+
|
|
119
|
+
memory_id = await self._s.scalar(
|
|
120
|
+
insert(Memory)
|
|
121
|
+
.values(
|
|
122
|
+
content=str(content),
|
|
123
|
+
memory_type=MemoryType(memory_type),
|
|
124
|
+
extra_data=metadata or {},
|
|
125
|
+
embedding=vector,
|
|
126
|
+
created_by=self._uid,
|
|
127
|
+
workspace_id=workspace_id,
|
|
128
|
+
)
|
|
129
|
+
.returning(Memory.id)
|
|
130
|
+
)
|
|
131
|
+
assert memory_id is not None
|
|
132
|
+
|
|
133
|
+
await self._upsert_tags(memory_id, tags or set())
|
|
134
|
+
return memory_id, True
|
|
135
|
+
|
|
136
|
+
async def update(
|
|
137
|
+
self,
|
|
138
|
+
id: int,
|
|
139
|
+
workspace_ids: list[int] | None = None,
|
|
140
|
+
content: str = _UNSET, # type: ignore[assignment]
|
|
141
|
+
memory_type: MemoryType = _UNSET, # type: ignore[assignment]
|
|
142
|
+
metadata: dict = _UNSET, # type: ignore[assignment]
|
|
143
|
+
tags: set[str] = _UNSET, # type: ignore[assignment]
|
|
144
|
+
) -> int:
|
|
145
|
+
"""Update memory fields. Pass _UNSET to leave a field unchanged; None sets it to NULL."""
|
|
146
|
+
if workspace_ids is None:
|
|
147
|
+
workspace_ids = await self._accessible_workspace_ids()
|
|
148
|
+
access_check = Memory.workspace_id.in_(workspace_ids)
|
|
149
|
+
|
|
150
|
+
memory_id = await self._s.scalar(
|
|
151
|
+
select(Memory.id).where(Memory.id == id, access_check)
|
|
152
|
+
)
|
|
153
|
+
if memory_id is None:
|
|
154
|
+
raise ValueError(f"Memory with id={id} not found")
|
|
155
|
+
|
|
156
|
+
values: dict = {}
|
|
157
|
+
if memory_type is not _UNSET:
|
|
158
|
+
values["memory_type"] = MemoryType(memory_type)
|
|
159
|
+
if metadata is not _UNSET:
|
|
160
|
+
values["extra_data"] = metadata or {}
|
|
161
|
+
|
|
162
|
+
if content is not _UNSET or tags is not _UNSET:
|
|
163
|
+
new_content = (
|
|
164
|
+
content
|
|
165
|
+
if content is not _UNSET
|
|
166
|
+
else (
|
|
167
|
+
await self._s.scalar(
|
|
168
|
+
select(Memory.content).where(Memory.id == memory_id)
|
|
169
|
+
)
|
|
170
|
+
or ""
|
|
171
|
+
)
|
|
172
|
+
)
|
|
173
|
+
new_tags = (
|
|
174
|
+
list(tags)
|
|
175
|
+
if tags is not _UNSET
|
|
176
|
+
else await self._fetch_tag_names(memory_id)
|
|
177
|
+
)
|
|
178
|
+
if content is not _UNSET:
|
|
179
|
+
values["content"] = content
|
|
180
|
+
values["embedding"] = await embed(_embed_text(new_content, new_tags))
|
|
181
|
+
|
|
182
|
+
if values:
|
|
183
|
+
await self._s.execute(
|
|
184
|
+
update(Memory).where(Memory.id == memory_id).values(**values)
|
|
185
|
+
)
|
|
186
|
+
|
|
187
|
+
if tags is not _UNSET:
|
|
188
|
+
await self._replace_tags(memory_id, tags) # type: ignore[arg-type]
|
|
189
|
+
|
|
190
|
+
return memory_id
|
|
191
|
+
|
|
192
|
+
async def delete(self, id: int, workspace_ids: list[int] | None = None) -> None:
|
|
193
|
+
if workspace_ids is None:
|
|
194
|
+
workspace_ids = await self._accessible_workspace_ids()
|
|
195
|
+
access_check = Memory.workspace_id.in_(workspace_ids)
|
|
196
|
+
|
|
197
|
+
result = await self._s.scalar(
|
|
198
|
+
delete(Memory).where(Memory.id == id, access_check).returning(Memory.id)
|
|
199
|
+
)
|
|
200
|
+
if result is None:
|
|
201
|
+
raise ValueError(f"Memory with id={id} not found")
|
|
202
|
+
await self._cleanup_orphan_tags()
|
|
203
|
+
|
|
204
|
+
async def get(
|
|
205
|
+
self, id: int, workspace_ids: list[int] | None = None
|
|
206
|
+
) -> MemoryListItem | None:
|
|
207
|
+
if workspace_ids is None:
|
|
208
|
+
workspace_ids = await self._accessible_workspace_ids()
|
|
209
|
+
access_check = Memory.workspace_id.in_(workspace_ids)
|
|
210
|
+
|
|
211
|
+
row = (
|
|
212
|
+
(
|
|
213
|
+
await self._s.execute(
|
|
214
|
+
select(
|
|
215
|
+
Memory.id,
|
|
216
|
+
Memory.content,
|
|
217
|
+
Memory.memory_type,
|
|
218
|
+
Memory.extra_data.label("metadata"),
|
|
219
|
+
Memory.created_at,
|
|
220
|
+
Memory.workspace_id,
|
|
221
|
+
).where(Memory.id == id, access_check)
|
|
222
|
+
)
|
|
223
|
+
)
|
|
224
|
+
.mappings()
|
|
225
|
+
.one_or_none()
|
|
226
|
+
)
|
|
227
|
+
if row is None:
|
|
228
|
+
return None
|
|
229
|
+
tags = (await self.fetch_tags([id])).get(id, [])
|
|
230
|
+
return MemoryListItem(
|
|
231
|
+
**row,
|
|
232
|
+
tags=tags,
|
|
233
|
+
)
|
|
234
|
+
|
|
235
|
+
async def move(self, id: int, target_workspace_id: int) -> None:
|
|
236
|
+
"""Move memory to a different workspace. Raises ValueError if not found or duplicate."""
|
|
237
|
+
workspace_ids = await self._accessible_workspace_ids()
|
|
238
|
+
row = (
|
|
239
|
+
(
|
|
240
|
+
await self._s.execute(
|
|
241
|
+
select(Memory.id, Memory.content).where(
|
|
242
|
+
Memory.id == id, Memory.workspace_id.in_(workspace_ids)
|
|
243
|
+
)
|
|
244
|
+
)
|
|
245
|
+
)
|
|
246
|
+
.mappings()
|
|
247
|
+
.one_or_none()
|
|
248
|
+
)
|
|
249
|
+
if row is None:
|
|
250
|
+
raise ValueError(f"Memory with id={id} not found")
|
|
251
|
+
if target_workspace_id not in workspace_ids:
|
|
252
|
+
raise PermissionError(
|
|
253
|
+
f"No access to target workspace {target_workspace_id}"
|
|
254
|
+
)
|
|
255
|
+
duplicate = await self._s.scalar(
|
|
256
|
+
select(Memory.id).where(
|
|
257
|
+
Memory.content == row["content"],
|
|
258
|
+
Memory.workspace_id == target_workspace_id,
|
|
259
|
+
Memory.id != id,
|
|
260
|
+
)
|
|
261
|
+
)
|
|
262
|
+
if duplicate is not None:
|
|
263
|
+
raise ValueError(
|
|
264
|
+
"A memory with the same content already exists in the target workspace"
|
|
265
|
+
)
|
|
266
|
+
await self._s.execute(
|
|
267
|
+
update(Memory)
|
|
268
|
+
.where(Memory.id == id)
|
|
269
|
+
.values(workspace_id=target_workspace_id)
|
|
270
|
+
)
|
|
271
|
+
|
|
272
|
+
async def fetch_tags(self, memory_ids: list[int]) -> dict[int, list[str]]:
|
|
273
|
+
rows = await self._s.execute(
|
|
274
|
+
select(MemoryTag.memory_id, Tag.name)
|
|
275
|
+
.join(Tag, MemoryTag.tag_id == Tag.id)
|
|
276
|
+
.where(MemoryTag.memory_id.in_(memory_ids))
|
|
277
|
+
)
|
|
278
|
+
result: dict[int, list[str]] = {i: [] for i in memory_ids}
|
|
279
|
+
for mid, name in rows.fetchall():
|
|
280
|
+
result[mid].append(name)
|
|
281
|
+
return result
|
|
282
|
+
|
|
283
|
+
async def fetch_metadata(self, memory_ids: list[int]) -> dict[int, tuple]:
|
|
284
|
+
rows = await self._s.execute(
|
|
285
|
+
select(Memory.id, Memory.extra_data, Memory.created_at).where(
|
|
286
|
+
Memory.id.in_(memory_ids)
|
|
287
|
+
)
|
|
288
|
+
)
|
|
289
|
+
return {row.id: (row.extra_data, row.created_at) for row in rows.fetchall()}
|
memlord/dao/user.py
ADDED
|
@@ -0,0 +1,105 @@
|
|
|
1
|
+
from sqlalchemy import insert, select, update
|
|
2
|
+
from sqlalchemy.ext.asyncio import AsyncSession
|
|
3
|
+
|
|
4
|
+
from memlord.auth import verify_password
|
|
5
|
+
from memlord.dao.workspace import WorkspaceDao
|
|
6
|
+
from memlord.models.user import User
|
|
7
|
+
from memlord.schemas.user import UserInfo
|
|
8
|
+
|
|
9
|
+
|
|
10
|
+
class UserDao:
|
|
11
|
+
def __init__(self, s: AsyncSession) -> None:
|
|
12
|
+
self._s = s
|
|
13
|
+
|
|
14
|
+
async def authenticate(self, email: str, password: str) -> UserInfo | None:
|
|
15
|
+
row = (
|
|
16
|
+
(
|
|
17
|
+
await self._s.execute(
|
|
18
|
+
select(
|
|
19
|
+
User.id,
|
|
20
|
+
User.display_name,
|
|
21
|
+
User.email,
|
|
22
|
+
User.email_verified,
|
|
23
|
+
User.hashed_password,
|
|
24
|
+
).where(User.email == email.strip().lower())
|
|
25
|
+
)
|
|
26
|
+
)
|
|
27
|
+
.mappings()
|
|
28
|
+
.one_or_none()
|
|
29
|
+
)
|
|
30
|
+
if row is None or not verify_password(password, row["hashed_password"]):
|
|
31
|
+
return None
|
|
32
|
+
return UserInfo(
|
|
33
|
+
id=row["id"],
|
|
34
|
+
display_name=row["display_name"],
|
|
35
|
+
email=row["email"],
|
|
36
|
+
email_verified=row["email_verified"],
|
|
37
|
+
)
|
|
38
|
+
|
|
39
|
+
async def exists_by_email(self, email: str) -> bool:
|
|
40
|
+
result = await self._s.scalar(
|
|
41
|
+
select(User.id).where(User.email == email.strip().lower())
|
|
42
|
+
)
|
|
43
|
+
return result is not None
|
|
44
|
+
|
|
45
|
+
async def get_by_id(self, id: int) -> UserInfo | None:
|
|
46
|
+
row = (
|
|
47
|
+
(
|
|
48
|
+
await self._s.execute(
|
|
49
|
+
select(
|
|
50
|
+
User.id, User.display_name, User.email, User.email_verified
|
|
51
|
+
).where(User.id == id)
|
|
52
|
+
)
|
|
53
|
+
)
|
|
54
|
+
.mappings()
|
|
55
|
+
.one_or_none()
|
|
56
|
+
)
|
|
57
|
+
if row is None:
|
|
58
|
+
return None
|
|
59
|
+
return UserInfo(
|
|
60
|
+
id=row["id"],
|
|
61
|
+
display_name=row["display_name"],
|
|
62
|
+
email=row["email"],
|
|
63
|
+
email_verified=row["email_verified"],
|
|
64
|
+
)
|
|
65
|
+
|
|
66
|
+
async def get_email_by_id(self, id: int) -> str | None:
|
|
67
|
+
return await self._s.scalar(select(User.email).where(User.id == id))
|
|
68
|
+
|
|
69
|
+
async def get_id_by_email(self, email: str) -> int | None:
|
|
70
|
+
return await self._s.scalar(
|
|
71
|
+
select(User.id).where(User.email == email.strip().lower())
|
|
72
|
+
)
|
|
73
|
+
|
|
74
|
+
async def set_email_verified(self, user_id: int) -> None:
|
|
75
|
+
await self._s.execute(
|
|
76
|
+
update(User).where(User.id == user_id).values(email_verified=True)
|
|
77
|
+
)
|
|
78
|
+
|
|
79
|
+
async def set_password(self, user_id: int, hashed_password: str) -> None:
|
|
80
|
+
await self._s.execute(
|
|
81
|
+
update(User)
|
|
82
|
+
.where(User.id == user_id)
|
|
83
|
+
.values(hashed_password=hashed_password)
|
|
84
|
+
)
|
|
85
|
+
|
|
86
|
+
async def create(
|
|
87
|
+
self, email: str, display_name: str, hashed_password: str
|
|
88
|
+
) -> UserInfo:
|
|
89
|
+
user_id = await self._s.scalar(
|
|
90
|
+
insert(User)
|
|
91
|
+
.values(
|
|
92
|
+
email=email.strip().lower(),
|
|
93
|
+
display_name=display_name.strip(),
|
|
94
|
+
hashed_password=hashed_password,
|
|
95
|
+
)
|
|
96
|
+
.returning(User.id)
|
|
97
|
+
)
|
|
98
|
+
assert user_id is not None
|
|
99
|
+
await WorkspaceDao(self._s).create_personal(user_id)
|
|
100
|
+
return UserInfo(
|
|
101
|
+
id=user_id,
|
|
102
|
+
display_name=display_name.strip(),
|
|
103
|
+
email=email.strip().lower(),
|
|
104
|
+
email_verified=False,
|
|
105
|
+
)
|