objbase 0.3.0__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.
@@ -0,0 +1,268 @@
1
+ import json
2
+ import os
3
+ import sys
4
+
5
+ from objbase.interface import InventoryStorage, Item
6
+ from objbase.util.file_util import atomic_write_json, atomic_write_text, locked
7
+
8
+ # Newlines are rejected because the directory storage index stores one id per line.
9
+ _UNSAFE_CHARS = "/\\\x00\n\r"
10
+ if sys.platform == "win32":
11
+ # Characters Windows doesn't allow in file names; ":" would also allow drive-relative
12
+ # paths ("C:x", which os.path.join doesn't anchor to the base dir) and alternate data streams.
13
+ _UNSAFE_CHARS += '<>:"|?*'
14
+
15
+
16
+ def _safe_name(name: str, kind: str) -> str:
17
+ """Validate that an item type or id can be used as a single path component."""
18
+ if not isinstance(name, str) or name in ("", ".", "..") or any(c in name for c in _UNSAFE_CHARS):
19
+ raise ValueError(f"Invalid {kind} for file storage: {name!r}")
20
+ if sys.platform == "win32" and name[-1] in ". ":
21
+ # Windows drops trailing dots and spaces, so "a." would alias "a" and ".. " would mean "..".
22
+ raise ValueError(f"Invalid {kind} for file storage: {name!r}")
23
+ return name
24
+
25
+
26
+ def _strip_extended_prefix(path: str) -> str:
27
+ """Remove a Windows extended-length prefix (``\\\\?\\`` or ``\\\\?\\UNC\\``) from ``path``.
28
+
29
+ ``os.path.realpath`` can leave the prefix on a path that doesn't exist yet when
30
+ its parent directory appears while it's resolving (another thread creating a
31
+ type directory), so the result wouldn't compare equal to the unprefixed base.
32
+ """
33
+ if sys.platform == "win32":
34
+ if path.startswith("\\\\?\\UNC\\"):
35
+ return "\\\\" + path[len("\\\\?\\UNC\\") :]
36
+ if path.startswith("\\\\?\\"):
37
+ return path[len("\\\\?\\") :]
38
+ return path
39
+
40
+
41
+ def _contained_path(real_base: str, *parts: str) -> str:
42
+ """Join ``parts`` onto ``real_base``, checking that the real path of the result lies inside ``real_base``.
43
+
44
+ ``real_base`` must already be a real path (``os.path.realpath``). Resolving
45
+ symlinks catches links inside the base directory that point outside it, which
46
+ the name checks in ``_safe_name`` can't see. Returns the joined (unresolved)
47
+ path, so replacing or deleting a symlinked file acts on the link itself.
48
+
49
+ The check can't stop an attacker who can write to the base directory and swaps
50
+ in a symlink between the check and the file operation.
51
+ """
52
+ path = os.path.join(real_base, *parts)
53
+ real_path = os.path.normcase(_strip_extended_prefix(os.path.realpath(path)))
54
+ base = os.path.normcase(_strip_extended_prefix(real_base))
55
+ try:
56
+ inside = real_path != base and os.path.commonpath([base, real_path]) == base
57
+ except ValueError: # on different drives (Windows)
58
+ inside = False
59
+ if not inside:
60
+ raise ValueError(f"Path for file storage resolves outside the base directory: {path!r}")
61
+ return path
62
+
63
+
64
+ class FileBasedInventoryStorage(InventoryStorage):
65
+ """Simple file-based storage that saves all items of a given inventory type in a single JSON file.
66
+
67
+ Safe for concurrent use by multiple threads and processes on the same machine:
68
+ writes hold an exclusive lock on ``.{item_type}.json.lock`` for the whole
69
+ read-modify-write cycle, and files are replaced atomically. Locks are advisory
70
+ and may not work on network file systems.
71
+ """
72
+
73
+ def __init__(self, base_dir: str):
74
+ self.inventory_dir = base_dir
75
+ if not os.path.exists(self.inventory_dir):
76
+ raise ValueError(f"Base directory {self.inventory_dir} does not exist.")
77
+ self._real_base = os.path.realpath(base_dir)
78
+
79
+ def keys(self, item_type: str) -> list[str]:
80
+ return [item["id"] for item in self.items(item_type)]
81
+
82
+ def items(self, item_type: str) -> list[Item]:
83
+ file_path = self._file_path(item_type)
84
+ if not os.path.exists(file_path):
85
+ return []
86
+ with locked(self._lock_path(item_type), shared=True):
87
+ return self._load(file_path)
88
+
89
+ def write(self, item_type: str, item: Item) -> bool:
90
+ file_path = self._file_path(item_type)
91
+ with locked(self._lock_path(item_type)):
92
+ items = self._load(file_path)
93
+ for i, existing_item in enumerate(items):
94
+ if existing_item["id"] == item["id"]:
95
+ items[i] = item
96
+ break
97
+ else:
98
+ items.append(item)
99
+ atomic_write_json(file_path, items)
100
+ return True
101
+
102
+ def read(self, item_type: str, id: str) -> Item | None:
103
+ items = self.items(item_type)
104
+ for item in items:
105
+ if item["id"] == id:
106
+ return item
107
+ return None
108
+
109
+ def delete(self, item_type: str, id: str) -> bool:
110
+ file_path = self._file_path(item_type)
111
+ if not os.path.exists(file_path):
112
+ return False
113
+ with locked(self._lock_path(item_type)):
114
+ items = self._load(file_path)
115
+ remaining = [item for item in items if item["id"] != id]
116
+ if len(remaining) == len(items):
117
+ return False
118
+ atomic_write_json(file_path, remaining)
119
+ return True
120
+
121
+ def _file_path(self, item_type: str) -> str:
122
+ return _contained_path(self._real_base, f"{_safe_name(item_type, 'item type')}.json")
123
+
124
+ def _lock_path(self, item_type: str) -> str:
125
+ return _contained_path(self._real_base, f".{_safe_name(item_type, 'item type')}.json.lock")
126
+
127
+ @staticmethod
128
+ def _load(file_path: str) -> list[Item]:
129
+ """Load a type file; the caller must hold its lock. A missing file means no items."""
130
+ try:
131
+ with open(file_path) as f:
132
+ items: list[Item] = json.load(f)
133
+ except FileNotFoundError:
134
+ return []
135
+ return items
136
+
137
+
138
+ class DirectoryBasedInventoryStorage(InventoryStorage):
139
+ """Alternative file-based storage that uses a directory per inventory type and individual files per item.
140
+
141
+ Each type directory also holds an index file (``.index``) listing the ids of
142
+ all items of that type, one per line, so ``keys()`` doesn't have to scan the
143
+ directory. Writes and deletes hold an exclusive lock on ``.index.lock`` while
144
+ they change the item file and the index, keeping both in step across threads
145
+ and processes. Locks are advisory and may not work on network file systems.
146
+ """
147
+
148
+ INDEX_FILE = ".index"
149
+
150
+ def __init__(self, base_dir: str):
151
+ self.inventory_dir = base_dir
152
+ if not os.path.exists(self.inventory_dir):
153
+ raise ValueError(f"Base directory {self.inventory_dir} does not exist.")
154
+ self._real_base = os.path.realpath(base_dir)
155
+
156
+ def _type_dir(self, item_type: str) -> str:
157
+ return _contained_path(self._real_base, _safe_name(item_type, "item type"))
158
+
159
+ def _item_path(self, item_type: str, id: str) -> str:
160
+ return self._path_in_type_dir(item_type, f"{_safe_name(id, 'item id')}.json")
161
+
162
+ def _path_in_type_dir(self, item_type: str, filename: str) -> str:
163
+ return _contained_path(self._real_base, _safe_name(item_type, "item type"), filename)
164
+
165
+ def _index_path(self, item_type: str) -> str:
166
+ return self._path_in_type_dir(item_type, self.INDEX_FILE)
167
+
168
+ def _lock_path(self, item_type: str) -> str:
169
+ return self._path_in_type_dir(item_type, f"{self.INDEX_FILE}.lock")
170
+
171
+ def keys(self, item_type: str) -> list[str]:
172
+ type_dir = self._type_dir(item_type)
173
+ if not os.path.exists(type_dir):
174
+ return []
175
+ with locked(self._lock_path(item_type), shared=True):
176
+ ids = self._read_index(item_type)
177
+ # A type directory without an index predates indexing; the next write or
178
+ # delete creates the index, until then the directory is scanned.
179
+ return ids if ids is not None else self._scan(type_dir)
180
+
181
+ def items(self, item_type: str) -> list[Item]:
182
+ type_dir = self._type_dir(item_type)
183
+ if not os.path.exists(type_dir):
184
+ return []
185
+ items = []
186
+ for filename in os.listdir(type_dir):
187
+ if filename.endswith(".json"):
188
+ try:
189
+ with open(self._path_in_type_dir(item_type, filename)) as f:
190
+ items.append(json.load(f))
191
+ except FileNotFoundError:
192
+ continue # deleted by another process since listdir()
193
+ return items
194
+
195
+ def write(self, item_type: str, item: Item) -> bool:
196
+ item_id = item.get("id")
197
+ if not item_id:
198
+ raise ValueError("Item must have an 'id' field.")
199
+ item_path = self._item_path(item_type, item_id)
200
+ os.makedirs(self._type_dir(item_type), exist_ok=True)
201
+ with locked(self._lock_path(item_type)):
202
+ atomic_write_json(item_path, item)
203
+ if item_id not in self._load_index(item_type):
204
+ with open(self._index_path(item_type), "a") as f:
205
+ f.write(f"{item_id}\n")
206
+ f.flush()
207
+ os.fsync(f.fileno())
208
+ return True
209
+
210
+ def read(self, item_type: str, id: str) -> Item | None:
211
+ item_path = self._item_path(item_type, id)
212
+ try:
213
+ with open(item_path) as f:
214
+ item: Item = json.load(f)
215
+ except FileNotFoundError:
216
+ return None
217
+ return item
218
+
219
+ def delete(self, item_type: str, id: str) -> bool:
220
+ item_path = self._item_path(item_type, id)
221
+ if not os.path.exists(self._type_dir(item_type)):
222
+ return False
223
+ with locked(self._lock_path(item_type)):
224
+ try:
225
+ os.remove(item_path)
226
+ except FileNotFoundError:
227
+ return False
228
+ ids = self._load_index(item_type)
229
+ if id in ids:
230
+ self._write_index(item_type, [i for i in ids if i != id])
231
+ return True
232
+
233
+ def rebuild_index(self, item_type: str) -> None:
234
+ """Recreate the index of ``item_type`` from the item files in its directory.
235
+
236
+ Only needed if the index got out of step with the item files, e.g. after a
237
+ crash between writing an item and updating the index, or after item files
238
+ were added or removed by hand.
239
+ """
240
+ type_dir = self._type_dir(item_type)
241
+ if not os.path.exists(type_dir):
242
+ return
243
+ with locked(self._lock_path(item_type)):
244
+ self._write_index(item_type, self._scan(type_dir))
245
+
246
+ @staticmethod
247
+ def _scan(type_dir: str) -> list[str]:
248
+ """Ids of all item files in ``type_dir``, from the directory listing (item files are named ``{id}.json``)."""
249
+ return [filename.removesuffix(".json") for filename in os.listdir(type_dir) if filename.endswith(".json")]
250
+
251
+ def _read_index(self, item_type: str) -> list[str] | None:
252
+ """Ids in the index of ``item_type``, or ``None`` if there is no index. The caller must hold the lock."""
253
+ try:
254
+ with open(self._index_path(item_type)) as f:
255
+ return [line for line in f.read().splitlines() if line]
256
+ except FileNotFoundError:
257
+ return None
258
+
259
+ def _load_index(self, item_type: str) -> list[str]:
260
+ """Like ``_read_index``, but creates a missing index first. The caller must hold the exclusive lock."""
261
+ ids = self._read_index(item_type)
262
+ if ids is None:
263
+ ids = self._scan(self._type_dir(item_type))
264
+ self._write_index(item_type, ids)
265
+ return ids
266
+
267
+ def _write_index(self, item_type: str, ids: list[str]) -> None:
268
+ atomic_write_text(self._index_path(item_type), "".join(f"{id}\n" for id in ids))
@@ -0,0 +1,51 @@
1
+ import copy
2
+
3
+ from objbase.asyncio.async_storage import AsyncInventoryStorage
4
+ from objbase.interface import InventoryStorage, Item
5
+
6
+
7
+ class InMemoryInventoryStorage(InventoryStorage, AsyncInventoryStorage):
8
+ """In-memory storage implementation for inventory items.
9
+
10
+ Items are deep-copied on the way in and out, so callers never share
11
+ mutable state with the store.
12
+ """
13
+
14
+ def __init__(self) -> None:
15
+ self.data: dict[str, dict[str, Item]] = {}
16
+
17
+ def keys(self, item_type: str) -> list[str]:
18
+ return list(self.data.get(item_type, {}))
19
+
20
+ def items(self, item_type: str) -> list[Item]:
21
+ return copy.deepcopy(list(self.data.get(item_type, {}).values()))
22
+
23
+ def read(self, item_type: str, id: str) -> Item | None:
24
+ return copy.deepcopy(self.data.get(item_type, {}).get(id))
25
+
26
+ def write(self, item_type: str, item: Item) -> bool:
27
+ if item_type not in self.data:
28
+ self.data[item_type] = {}
29
+ self.data[item_type][item["id"]] = copy.deepcopy(item)
30
+ return True
31
+
32
+ def delete(self, item_type: str, id: str) -> bool:
33
+ if item_type in self.data and id in self.data[item_type]:
34
+ del self.data[item_type][id]
35
+ return True
36
+ return False
37
+
38
+ async def akeys(self, item_type: str) -> list[str]:
39
+ return self.keys(item_type)
40
+
41
+ async def aitems(self, item_type: str) -> list[Item]:
42
+ return self.items(item_type)
43
+
44
+ async def aread(self, item_type: str, id: str) -> Item | None:
45
+ return self.read(item_type, id)
46
+
47
+ async def awrite(self, item_type: str, item: Item) -> bool:
48
+ return self.write(item_type, item)
49
+
50
+ async def adelete(self, item_type: str, id: str) -> bool:
51
+ return self.delete(item_type, id)
@@ -0,0 +1,42 @@
1
+ from collections.abc import Mapping
2
+ from typing import TYPE_CHECKING, Any
3
+
4
+ from objbase.interface import InventoryStorage, Item
5
+
6
+ if TYPE_CHECKING:
7
+ from pymongo import MongoClient
8
+ from pymongo.collection import Collection
9
+
10
+
11
+ class MongoDBInventoryStorage(InventoryStorage):
12
+ """MongoDB-based storage implementation for inventory items."""
13
+
14
+ def __init__(self, mongo_client: "MongoClient[Item]"):
15
+ self.mongo_client = mongo_client
16
+
17
+ def get_mongo_collection(self, item_type: str) -> "Collection[Item]":
18
+ db = self.mongo_client["inventory"]
19
+ return db[item_type]
20
+
21
+ def keys(self, item_type: str) -> list[str]:
22
+ collection = self.get_mongo_collection(item_type)
23
+ return [doc["id"] for doc in collection.find({}, {"id": True, "_id": False})]
24
+
25
+ def items(self, item_type: str, query: Mapping[str, Any] | None = None) -> list[Item]:
26
+ """Return all items of a type. ``query`` is a MongoDB-only extension to filter results."""
27
+ collection = self.get_mongo_collection(item_type)
28
+ return list(collection.find(query or {}, {"_id": False}))
29
+
30
+ def write(self, item_type: str, item: Item) -> bool:
31
+ collection = self.get_mongo_collection(item_type)
32
+ collection.replace_one({"id": item["id"]}, item, upsert=True)
33
+ return True
34
+
35
+ def read(self, item_type: str, id: str) -> Item | None:
36
+ collection = self.get_mongo_collection(item_type)
37
+ return collection.find_one({"id": id}, {"_id": False})
38
+
39
+ def delete(self, item_type: str, id: str) -> bool:
40
+ collection = self.get_mongo_collection(item_type)
41
+ result = collection.delete_one({"id": id})
42
+ return result.deleted_count > 0
@@ -0,0 +1,68 @@
1
+ import json
2
+ from typing import Any, Protocol
3
+
4
+ from objbase.interface import InventoryStorage, Item
5
+
6
+ DEFAULT_KEY_PREFIX = "inventory:"
7
+
8
+
9
+ class RedisHashClient(Protocol):
10
+ """The Redis hash commands the Redis adapters use.
11
+
12
+ Satisfied by ``redis.Redis`` and ``redis.asyncio.Redis`` (and compatible
13
+ clients). Return types are ``Any`` because redis-py annotates every command
14
+ as returning ``Awaitable[T] | T``; the async adapter awaits the results.
15
+ """
16
+
17
+ def hkeys(self, name: str, /) -> Any: ...
18
+
19
+ def hvals(self, name: str, /) -> Any: ...
20
+
21
+ def hget(self, name: str, key: str, /) -> Any: ...
22
+
23
+ def hset(self, name: str, key: str, value: str, /) -> Any: ...
24
+
25
+ def hdel(self, name: str, /, *keys: str) -> Any: ...
26
+
27
+
28
+ def decode_key(key: bytes | str) -> str:
29
+ """Hash field names are bytes unless the client was created with ``decode_responses=True``."""
30
+ return key.decode() if isinstance(key, bytes) else key
31
+
32
+
33
+ def redis_type_key(key_prefix: str, item_type: str) -> str:
34
+ """Name of the Redis hash that holds all items of ``item_type``, keyed by id."""
35
+ return f"{key_prefix}{item_type}"
36
+
37
+
38
+ class RedisInventoryStorage(InventoryStorage):
39
+ """Redis-backed storage.
40
+
41
+ Each item type is one Redis hash (``{key_prefix}{item_type}``) mapping item
42
+ ids to JSON-encoded items. Works with clients created with or without
43
+ ``decode_responses=True``.
44
+ """
45
+
46
+ def __init__(self, redis_client: RedisHashClient, key_prefix: str = DEFAULT_KEY_PREFIX):
47
+ self.redis_client = redis_client
48
+ self.key_prefix = key_prefix
49
+
50
+ def _key(self, item_type: str) -> str:
51
+ return redis_type_key(self.key_prefix, item_type)
52
+
53
+ def keys(self, item_type: str) -> list[str]:
54
+ return [decode_key(key) for key in self.redis_client.hkeys(self._key(item_type))]
55
+
56
+ def items(self, item_type: str) -> list[Item]:
57
+ return [json.loads(value) for value in self.redis_client.hvals(self._key(item_type))]
58
+
59
+ def write(self, item_type: str, item: Item) -> bool:
60
+ self.redis_client.hset(self._key(item_type), item["id"], json.dumps(item))
61
+ return True
62
+
63
+ def read(self, item_type: str, id: str) -> Item | None:
64
+ value = self.redis_client.hget(self._key(item_type), id)
65
+ return json.loads(value) if value is not None else None
66
+
67
+ def delete(self, item_type: str, id: str) -> bool:
68
+ return bool(self.redis_client.hdel(self._key(item_type), id))
@@ -0,0 +1,79 @@
1
+ import json
2
+ import sqlite3
3
+ from collections.abc import Iterator
4
+ from contextlib import contextmanager
5
+
6
+ from objbase.interface import InventoryStorage, Item
7
+
8
+ CREATE_TABLE_SQL = """
9
+ CREATE TABLE IF NOT EXISTS items (
10
+ item_type TEXT NOT NULL,
11
+ id TEXT NOT NULL,
12
+ data TEXT NOT NULL,
13
+ PRIMARY KEY (item_type, id)
14
+ )
15
+ """
16
+
17
+
18
+ class SQLiteInventoryStorage(InventoryStorage):
19
+ """SQLite-backed storage. Each item is stored as a JSON blob in a single table."""
20
+
21
+ def __init__(self, db_path: str):
22
+ self.db_path = db_path
23
+ with self._connect() as conn:
24
+ conn.execute(CREATE_TABLE_SQL)
25
+
26
+ @contextmanager
27
+ def _connect(self) -> Iterator[sqlite3.Connection]:
28
+ """Open a connection, commit on success (rollback on error), and always close it."""
29
+ conn = sqlite3.connect(self.db_path)
30
+ conn.row_factory = sqlite3.Row
31
+ try:
32
+ with conn:
33
+ yield conn
34
+ finally:
35
+ conn.close()
36
+
37
+ def keys(self, item_type: str) -> list[str]:
38
+ with self._connect() as conn:
39
+ rows = conn.execute(
40
+ "SELECT id FROM items WHERE item_type = ?",
41
+ (item_type,),
42
+ ).fetchall()
43
+ return [row["id"] for row in rows]
44
+
45
+ def items(self, item_type: str) -> list[Item]:
46
+ with self._connect() as conn:
47
+ rows = conn.execute(
48
+ "SELECT data FROM items WHERE item_type = ?",
49
+ (item_type,),
50
+ ).fetchall()
51
+ return [json.loads(row["data"]) for row in rows]
52
+
53
+ def read(self, item_type: str, id: str) -> Item | None:
54
+ with self._connect() as conn:
55
+ row = conn.execute(
56
+ "SELECT data FROM items WHERE item_type = ? AND id = ?",
57
+ (item_type, id),
58
+ ).fetchone()
59
+ return json.loads(row["data"]) if row else None
60
+
61
+ def write(self, item_type: str, item: Item) -> bool:
62
+ with self._connect() as conn:
63
+ conn.execute(
64
+ """
65
+ INSERT INTO items (item_type, id, data)
66
+ VALUES (?, ?, ?)
67
+ ON CONFLICT (item_type, id) DO UPDATE SET data = excluded.data
68
+ """,
69
+ (item_type, item["id"], json.dumps(item)),
70
+ )
71
+ return True
72
+
73
+ def delete(self, item_type: str, id: str) -> bool:
74
+ with self._connect() as conn:
75
+ cursor = conn.execute(
76
+ "DELETE FROM items WHERE item_type = ? AND id = ?",
77
+ (item_type, id),
78
+ )
79
+ return cursor.rowcount > 0
File without changes
@@ -0,0 +1,102 @@
1
+ """Cross-platform file locking and atomic JSON writes, using only the standard library."""
2
+
3
+ import contextlib
4
+ import json
5
+ import os
6
+ import shutil
7
+ import sys
8
+ import time
9
+ import uuid
10
+ from collections.abc import Iterator
11
+ from typing import IO, Any
12
+
13
+ if sys.platform == "win32":
14
+ import msvcrt
15
+
16
+ def _lock(f: IO[bytes], shared: bool) -> None:
17
+ # msvcrt has no shared locks, so readers lock exclusively too. LK_LOCK
18
+ # gives up after ~10s, so retry until the lock is acquired.
19
+ f.seek(0)
20
+ while True:
21
+ try:
22
+ msvcrt.locking(f.fileno(), msvcrt.LK_LOCK, 1)
23
+ return
24
+ except OSError:
25
+ time.sleep(0.05)
26
+
27
+ def _unlock(f: IO[bytes]) -> None:
28
+ f.seek(0)
29
+ msvcrt.locking(f.fileno(), msvcrt.LK_UNLCK, 1)
30
+
31
+ else:
32
+ import fcntl
33
+
34
+ def _lock(f: IO[bytes], shared: bool) -> None:
35
+ # flock locks belong to the open file, so they also exclude other threads
36
+ # in the same process, not just other processes.
37
+ fcntl.flock(f.fileno(), fcntl.LOCK_SH if shared else fcntl.LOCK_EX)
38
+
39
+ def _unlock(f: IO[bytes]) -> None:
40
+ fcntl.flock(f.fileno(), fcntl.LOCK_UN)
41
+
42
+
43
+ @contextlib.contextmanager
44
+ def locked(lock_path: str, shared: bool = False) -> Iterator[None]:
45
+ """Hold an advisory lock on ``lock_path`` (created if missing) for the duration of the block.
46
+
47
+ Blocks until the lock is available. ``shared=True`` allows concurrent readers
48
+ (on Windows all locks are exclusive). The lock is released if the process dies.
49
+ Lock files are left in place; deleting them while others may use them is unsafe.
50
+ """
51
+ with open(lock_path, "a+b") as f:
52
+ _lock(f, shared)
53
+ try:
54
+ yield
55
+ finally:
56
+ _unlock(f)
57
+
58
+
59
+ def _replace(src: str, dst: str) -> None:
60
+ if sys.platform != "win32":
61
+ os.replace(src, dst)
62
+ return
63
+ # Windows refuses to replace a file another process has open (e.g. an unlocked
64
+ # reader), so retry briefly before giving up.
65
+ for _ in range(40):
66
+ try:
67
+ os.replace(src, dst)
68
+ return
69
+ except PermissionError:
70
+ time.sleep(0.05)
71
+ os.replace(src, dst)
72
+
73
+
74
+ def atomic_write_json(path: str, data: Any) -> None:
75
+ """Write ``data`` as JSON to ``path`` so readers see either the old or the new file, never a partial one.
76
+
77
+ See ``atomic_write_text``.
78
+ """
79
+ atomic_write_text(path, json.dumps(data, indent=4))
80
+
81
+
82
+ def atomic_write_text(path: str, text: str) -> None:
83
+ """Write ``text`` to ``path`` so readers see either the old or the new file, never a partial one.
84
+
85
+ Writes to a temporary file in the same directory and renames it over ``path``.
86
+ Keeps the permissions of an existing ``path``; a new file gets the default
87
+ permissions for the current umask.
88
+ """
89
+ directory, name = os.path.split(path)
90
+ tmp_path = os.path.join(directory, f".{name}.{uuid.uuid4().hex}.tmp")
91
+ try:
92
+ with open(tmp_path, "x") as f:
93
+ f.write(text)
94
+ f.flush()
95
+ os.fsync(f.fileno())
96
+ if os.path.exists(path):
97
+ shutil.copymode(path, tmp_path)
98
+ _replace(tmp_path, path)
99
+ except BaseException:
100
+ with contextlib.suppress(FileNotFoundError):
101
+ os.remove(tmp_path)
102
+ raise
@@ -0,0 +1,45 @@
1
+ import os
2
+ from typing import TYPE_CHECKING, Any
3
+
4
+ if TYPE_CHECKING:
5
+ from pymongo import MongoClient
6
+
7
+
8
+ def get_mongo_client(uri: str | None = None, ping: bool = False) -> "MongoClient[dict[str, Any]]":
9
+ mongodb_uri = uri or os.getenv("MONGODB_URI")
10
+ if not mongodb_uri:
11
+ raise ValueError("MONGODB_URI is not set in environment variables.")
12
+
13
+ try:
14
+ import pymongo
15
+ except ImportError:
16
+ raise ImportError("pymongo is not installed. Please install it with 'pip install pymongo'.") from None
17
+ client: MongoClient[dict[str, Any]] = pymongo.MongoClient(mongodb_uri)
18
+
19
+ if ping:
20
+ try:
21
+ # The ping command is cheap and does not require auth.
22
+ client.admin.command("ping")
23
+ except Exception as e:
24
+ raise ConnectionError(f"Could not connect to MongoDB: {e}") from e
25
+ return client
26
+
27
+
28
+ def mongodb_results_to_json(results: list[dict[str, Any]], strip_id: bool = True) -> list[dict[str, Any]]:
29
+ json_results = []
30
+ for doc in results:
31
+ if strip_id and "_id" in doc:
32
+ doc.pop("_id")
33
+ elif "_id" in doc:
34
+ doc["_id"] = str(doc["_id"]) # Convert ObjectId to string
35
+ json_results.append(doc)
36
+ return json_results
37
+
38
+
39
+ def mongodb_result_to_json(result: dict[str, Any], strip_id: bool = True) -> dict[str, Any]:
40
+ if result and "_id" in result:
41
+ if strip_id and "_id" in result:
42
+ result.pop("_id")
43
+ elif "_id" in result:
44
+ result["_id"] = str(result["_id"]) # Convert ObjectId to string
45
+ return result