feat(sync): 实现设备认证与对象版本协议原型

This commit is contained in:
2026-09-07 15:45:37 +08:00
parent f193f699b2
commit 7fffbcd55a
17 changed files with 1444 additions and 0 deletions
+1
View File
@@ -0,0 +1 @@
"""独立同步服务;不加载 AI Core、模型或本地 Vault。"""
+32
View File
@@ -0,0 +1,32 @@
"""运维入口通过终端隐式输入密码,不接受命令行秘密。"""
import argparse
import getpass
import os
from pathlib import Path
from .app import create_app
from .database import Database
from .storage import S3Objects
def main():
parser = argparse.ArgumentParser()
parser.add_argument("command", choices=["serve", "migrate", "create-user"])
parser.add_argument("--username")
args = parser.parse_args()
url = os.environ["SYNC_DATABASE_URL"]
if not url.startswith("postgresql+psycopg://"):
raise SystemExit("生产入口只支持 PostgreSQL")
db = Database(url)
db.migrate()
if args.command == "create-user":
db.add_user(args.username or input("用户名: "), getpass.getpass("密码(至少12字符): "))
elif args.command == "serve":
import uvicorn
app = create_app(db, S3Objects(os.environ["SYNC_S3_ENDPOINT"], os.environ["SYNC_S3_BUCKET"]), Path(os.environ["SYNC_STAGING_DIR"]))
uvicorn.run(app, host="0.0.0.0", port=8080, access_log=False)
if __name__ == "__main__":
main()
+296
View File
@@ -0,0 +1,296 @@
"""Sync v1 HTTP 边界;每次访问重新检查设备撤销,内容不进入日志。"""
import hashlib
import json
import secrets
import time
from pathlib import Path
from fastapi import FastAPI, Header, Query, Request
from fastapi.exceptions import RequestValidationError
from fastapi.responses import JSONResponse, Response
from .database import Database, password_hash, row, rows, run
from .models import Commit, Login, Refresh, Upload, VaultCreate
class SyncError(Exception):
def __init__(self, status, code, details=None):
self.status, self.code, self.details = status, code, details or {}
def digest(value: str) -> str:
return hashlib.sha256(value.encode()).hexdigest()
def create_app(db: Database, objects, staging: Path, *, quota=1024**3, clock=time.time):
db.migrate()
staging.mkdir(parents=True, exist_ok=True)
app = FastAPI(title="NotesAgent Sync", version="1.0.0")
app.state.database = db
@app.exception_handler(SyncError)
async def error(_request, exc):
return JSONResponse({"error": {"code": exc.code, "details": exc.details}}, status_code=exc.status)
@app.exception_handler(RequestValidationError)
async def invalid(_request, _exc):
# Pydantic 的原始错误可能带请求正文,禁止回显密码或笔记。
return JSONResponse({"error": {"code": "INVALID_REQUEST", "details": {}}}, status_code=422)
def identity(conn, authorization):
token = authorization.removeprefix("Bearer ") if authorization.startswith("Bearer ") else ""
session = row(conn, "SELECT s.*, d.user_id, d.revoked FROM sessions s JOIN devices d ON d.id=s.device_id WHERE s.token=:token", token=digest(token))
if not session or session["revoked"] or session["expires"] <= clock():
raise SyncError(401, "SESSION_EXPIRED")
return session
def vault(conn, vault_id, authorization, *, lock=False):
session = identity(conn, authorization)
# Vault 行锁覆盖配额、CAS、路径冲突和序列分配,跨 Worker 也保持一致。
suffix = " FOR UPDATE" if lock and not db.sqlite else ""
item = row(conn, "SELECT * FROM vaults WHERE id=:id AND user_id=:owner" + suffix,
id=vault_id, owner=session["user_id"])
if not item:
raise SyncError(404, "VAULT_NOT_FOUND")
return session, item
def issue(conn, device_id):
access, refresh = secrets.token_urlsafe(32), secrets.token_urlsafe(48)
run(conn, "INSERT INTO sessions VALUES (:token,:refresh,:device,:expires,:refresh_expires)",
token=digest(access), refresh=digest(refresh), device=device_id,
expires=int(clock()) + 900, refresh_expires=int(clock()) + 30 * 86400)
return {"access_token": access, "refresh_token": refresh, "expires_in": 900, "device_id": device_id}
@app.get("/health")
def health():
return {"status": "ok"}
@app.get("/ready")
def ready():
with db.transaction() as conn:
if row(conn, "SELECT version FROM schema_version")["version"] != 1:
raise SyncError(503, "SCHEMA_INCOMPATIBLE")
return {"status": "ready", "schema": 1}
@app.get("/sync/v1/handshake")
def handshake(protocol: int = 1):
if protocol != 1:
raise SyncError(426, "PROTOCOL_INCOMPATIBLE")
return {"protocol": 1, "max_object_size": 104857600, "chunk_size": 1048576,
"encryption": "transport-only", "history_retention": "indefinite",
"cursor_retention": "indefinite", "sharing": False}
@app.post("/sync/v1/auth/sessions")
def login(body: Login, request: Request):
key = digest((request.client.host if request.client else "unknown") + ":" + body.username.casefold())
# 失败计数先独立提交,抛出认证异常也不会回滚限流状态。
with db.transaction() as conn:
limit = row(conn, "SELECT * FROM login_limits WHERE key=:key", key=key)
if limit and clock() - limit["started"] < 60:
if limit["attempts"] >= 10:
raise SyncError(429, "RATE_LIMITED")
run(conn, "UPDATE login_limits SET attempts=attempts+1 WHERE key=:key", key=key)
else:
run(conn, "DELETE FROM login_limits WHERE key=:key", key=key)
run(conn, "INSERT INTO login_limits VALUES (:key,:now,1)", key=key, now=int(clock()))
with db.transaction() as conn:
user = row(conn, "SELECT * FROM users WHERE username=:name", name=body.username)
expected = user["password"] if user else password_hash("unavailable-user-password", "0" * 32)
if not secrets.compare_digest(password_hash(body.password, expected.split(":")[0]), expected) or not user:
raise SyncError(401, "LOGIN_FAILED")
device_id = secrets.token_hex(16)
run(conn, "INSERT INTO devices VALUES (:id,:user,:name,0)", id=device_id, user=user["id"], name=body.device_name)
return issue(conn, device_id)
@app.post("/sync/v1/auth/refresh")
def refresh(body: Refresh):
with db.transaction() as conn:
suffix = " FOR UPDATE OF s" if not db.sqlite else ""
session = row(conn, "SELECT s.*,d.revoked FROM sessions s JOIN devices d ON d.id=s.device_id WHERE refresh=:refresh" + suffix,
refresh=digest(body.refresh_token))
if not session or session["revoked"] or session["refresh_expires"] <= clock():
raise SyncError(401, "SESSION_EXPIRED")
run(conn, "DELETE FROM sessions WHERE token=:token", token=session["token"])
return issue(conn, session["device_id"])
@app.delete("/sync/v1/auth/sessions", status_code=204)
def logout(authorization: str = Header(default="")):
with db.transaction() as conn:
session = identity(conn, authorization)
run(conn, "DELETE FROM sessions WHERE token=:token", token=session["token"])
@app.get("/sync/v1/devices")
def devices(authorization: str = Header(default="")):
with db.transaction() as conn:
session = identity(conn, authorization)
return {"items": rows(conn, "SELECT id,name,revoked FROM devices WHERE user_id=:user", user=session["user_id"])}
@app.delete("/sync/v1/devices/{device_id}", status_code=204)
def revoke(device_id: str, authorization: str = Header(default="")):
with db.transaction() as conn:
session = identity(conn, authorization)
run(conn, "UPDATE devices SET revoked=1 WHERE id=:id AND user_id=:user", id=device_id, user=session["user_id"])
@app.post("/sync/v1/vaults")
def create_vault(body: VaultCreate, authorization: str = Header(default="")):
with db.transaction() as conn:
session = identity(conn, authorization)
vault_id = secrets.token_hex(16)
run(conn, "INSERT INTO vaults VALUES (:id,:user,:name,0,:quota,0)", id=vault_id, user=session["user_id"], name=body.name, quota=quota)
return {"vault_id": vault_id, "name": body.name}
@app.get("/sync/v1/vaults")
def list_vaults(authorization: str = Header(default="")):
with db.transaction() as conn:
session = identity(conn, authorization)
return {"items": rows(conn, "SELECT id,name,sequence,used,quota FROM vaults WHERE user_id=:user", user=session["user_id"])}
@app.post("/sync/v1/vaults/{vault_id}/uploads")
def begin_upload(vault_id: str, body: Upload, authorization: str = Header(default="")):
with db.transaction() as conn:
session, item = vault(conn, vault_id, authorization, lock=True)
found = row(conn, "SELECT * FROM objects WHERE vault_id=:v AND hash=:h", v=vault_id, h=body.content_hash)
if found:
if found["size"] != body.size:
raise SyncError(409, "OBJECT_SIZE_MISMATCH")
return {"complete": True, "upload_id": None, "offset": body.size}
reserved = row(conn, "SELECT COALESCE(SUM(size),0) AS size FROM uploads WHERE vault_id=:v AND expires>:now", v=vault_id, now=int(clock()))["size"]
if item["used"] + reserved + body.size > item["quota"]:
raise SyncError(413, "QUOTA_EXCEEDED")
upload_id = secrets.token_hex(16)
run(conn, "INSERT INTO uploads VALUES (:id,:v,:device,:h,:size,0,:expires)", id=upload_id,
v=vault_id, device=session["device_id"], h=body.content_hash, size=body.size, expires=int(clock()) + 3600)
(staging / upload_id).write_bytes(b"")
return {"complete": False, "upload_id": upload_id, "offset": 0}
def authorized_upload(conn, vault_id, upload_id, authorization):
session, _ = vault(conn, vault_id, authorization, lock=True)
upload = row(conn, "SELECT * FROM uploads WHERE id=:id AND vault_id=:v AND device_id=:device", id=upload_id, v=vault_id, device=session["device_id"])
if not upload or upload["expires"] <= clock():
raise SyncError(404, "UPLOAD_EXPIRED")
return upload
@app.get("/sync/v1/vaults/{vault_id}/uploads/{upload_id}")
def upload_status(vault_id: str, upload_id: str, authorization: str = Header(default="")):
with db.transaction() as conn:
upload = authorized_upload(conn, vault_id, upload_id, authorization)
return {"offset": upload["offset_bytes"], "size": upload["size"], "expires": upload["expires"]}
@app.put("/sync/v1/vaults/{vault_id}/uploads/{upload_id}")
async def upload_chunk(vault_id: str, upload_id: str, request: Request,
offset: int = Query(ge=0), authorization: str = Header(default="")):
data = bytearray()
async for chunk in request.stream():
data.extend(chunk)
if len(data) > 1048576:
raise SyncError(413, "CHUNK_TOO_LARGE")
with db.transaction() as conn:
upload = authorized_upload(conn, vault_id, upload_id, authorization)
if offset != upload["offset_bytes"]:
raise SyncError(409, "UPLOAD_OFFSET", {"offset": upload["offset_bytes"]})
if offset + len(data) > upload["size"]:
raise SyncError(413, "OBJECT_TOO_LARGE")
# 先刷盘后提交 offset;崩溃重试覆盖未确认尾部,不能重复追加。
import os
with (staging / upload_id).open("r+b") as stream:
stream.seek(offset)
stream.write(data)
stream.truncate()
stream.flush()
os.fsync(stream.fileno())
run(conn, "UPDATE uploads SET offset_bytes=:offset WHERE id=:id", id=upload_id, offset=offset + len(data))
return {"offset": offset + len(data)}
@app.delete("/sync/v1/vaults/{vault_id}/uploads/{upload_id}", status_code=204)
def cancel_upload(vault_id: str, upload_id: str, authorization: str = Header(default="")):
with db.transaction() as conn:
authorized_upload(conn, vault_id, upload_id, authorization)
run(conn, "DELETE FROM uploads WHERE id=:id", id=upload_id)
(staging / upload_id).unlink(missing_ok=True)
@app.post("/sync/v1/vaults/{vault_id}/uploads/{upload_id}/complete")
def complete_upload(vault_id: str, upload_id: str, authorization: str = Header(default="")):
with db.transaction() as conn:
upload = authorized_upload(conn, vault_id, upload_id, authorization)
data = (staging / upload_id).read_bytes()
if len(data) != upload["size"] or len(data) != upload["offset_bytes"] or hashlib.sha256(data).hexdigest() != upload["hash"]:
raise SyncError(422, "OBJECT_INTEGRITY")
exists = row(conn, "SELECT hash FROM objects WHERE vault_id=:v AND hash=:h", v=vault_id, h=upload["hash"])
if not exists:
objects.put(vault_id + "/" + upload["hash"], data)
run(conn, "INSERT INTO objects VALUES (:v,:h,:size,:now)", v=vault_id, h=upload["hash"], size=len(data), now=int(clock()))
run(conn, "UPDATE vaults SET used=used+:size WHERE id=:v", size=len(data), v=vault_id)
run(conn, "DELETE FROM uploads WHERE id=:id", id=upload_id)
(staging / upload_id).unlink(missing_ok=True)
return {"complete": True, "content_hash": upload["hash"]}
def perform_commit(conn, vault_id, body, session, item):
fingerprint = digest(body.model_dump_json())
previous = row(conn, "SELECT * FROM revisions WHERE vault_id=:v AND operation_id=:op", v=vault_id, op=body.operation_id)
if previous:
if previous["fingerprint"] != fingerprint or previous["device_id"] != session["device_id"]:
raise SyncError(409, "IDEMPOTENCY_REUSED")
return dict(previous)
current = row(conn, "SELECT * FROM files WHERE vault_id=:v AND file_id=:f", v=vault_id, f=body.file_id)
if (current["sequence"] if current else 0) != body.base_revision:
actual = row(conn, "SELECT * FROM revisions WHERE vault_id=:v AND sequence=:s", v=vault_id, s=current["sequence"]) if current else None
raise SyncError(409, "REVISION_CONFLICT", {"current": dict(actual) if actual else None})
if body.operation == "put":
obj = row(conn, "SELECT * FROM objects WHERE vault_id=:v AND hash=:h", v=vault_id, h=body.content_hash)
if not obj or obj["size"] != body.size:
raise SyncError(409, "OBJECT_NOT_READY")
paths = rows(conn, "SELECT path_key FROM files WHERE vault_id=:v AND deleted=0 AND file_id<>:f", v=vault_id, f=body.file_id)
key = body.path.casefold()
if any(p["path_key"] == key or p["path_key"].startswith(key + "/") or key.startswith(p["path_key"] + "/") for p in paths):
raise SyncError(409, "PATH_CONFLICT")
elif not current or body.content_hash is not None or body.size != 0:
raise SyncError(422, "INVALID_DELETE")
sequence = item["sequence"] + 1
run(conn, "UPDATE vaults SET sequence=:s WHERE id=:v", s=sequence, v=vault_id)
run(conn, "INSERT INTO revisions VALUES (:v,:s,:f,:base,:path,:key,:operation,:hash,:size,:device,:op,:fingerprint)",
v=vault_id, s=sequence, f=body.file_id, base=body.base_revision, path=body.path, key=body.path.casefold(),
operation=body.operation, hash=body.content_hash, size=body.size, device=session["device_id"], op=body.operation_id, fingerprint=fingerprint)
run(conn, "DELETE FROM files WHERE vault_id=:v AND file_id=:f", v=vault_id, f=body.file_id)
run(conn, "INSERT INTO files VALUES (:v,:f,:s,:key,:deleted)", v=vault_id, f=body.file_id, s=sequence,
key=body.path.casefold(), deleted=int(body.operation == "delete"))
return dict(row(conn, "SELECT * FROM revisions WHERE vault_id=:v AND sequence=:s", v=vault_id, s=sequence))
@app.post("/sync/v1/vaults/{vault_id}/revisions")
def commit(vault_id: str, body: Commit, authorization: str = Header(default="")):
with db.transaction() as conn:
session, item = vault(conn, vault_id, authorization, lock=True)
return perform_commit(conn, vault_id, body, session, item)
@app.get("/sync/v1/vaults/{vault_id}/changes")
def changes(vault_id: str, cursor: int = Query(default=0, ge=0), limit: int = Query(default=100, ge=1, le=500),
boundary: int | None = Query(default=None, ge=0), authorization: str = Header(default="")):
with db.transaction() as conn:
_, item = vault(conn, vault_id, authorization)
end = item["sequence"] if boundary is None else boundary
if cursor > end or end > item["sequence"]:
raise SyncError(409, "CURSOR_INVALID")
items = rows(conn, "SELECT * FROM revisions WHERE vault_id=:v AND sequence>:cursor AND sequence<=:end ORDER BY sequence LIMIT :limit", v=vault_id, cursor=cursor, end=end, limit=limit)
next_cursor = items[-1]["sequence"] if items else cursor
return {"items": items, "cursor": next_cursor, "boundary": end, "has_more": next_cursor < end}
@app.get("/sync/v1/vaults/{vault_id}/history/{file_id}")
def history(vault_id: str, file_id: str, before: int = Query(default=9223372036854775807, ge=1),
limit: int = Query(default=100, ge=1, le=500), authorization: str = Header(default="")):
with db.transaction() as conn:
vault(conn, vault_id, authorization)
return {"items": rows(conn, "SELECT * FROM revisions WHERE vault_id=:v AND file_id=:f AND sequence<:before ORDER BY sequence DESC LIMIT :limit", v=vault_id, f=file_id, before=before, limit=limit)}
@app.get("/sync/v1/vaults/{vault_id}/objects/{content_hash}")
def get_object(vault_id: str, content_hash: str, authorization: str = Header(default="")):
with db.transaction() as conn:
vault(conn, vault_id, authorization)
obj = row(conn, "SELECT * FROM objects WHERE vault_id=:v AND hash=:h", v=vault_id, h=content_hash)
if not obj:
raise SyncError(404, "OBJECT_NOT_FOUND")
data = objects.get(vault_id + "/" + content_hash)
if len(data) != obj["size"] or hashlib.sha256(data).hexdigest() != content_hash:
raise SyncError(503, "STORAGE_INTEGRITY")
return Response(data, media_type="application/octet-stream", headers={"ETag": '"' + content_hash + '"', "Cache-Control": "private, no-store"})
return app
+73
View File
@@ -0,0 +1,73 @@
"""Schema v1 与事务边界;生产使用 PostgreSQL,SQLite 仅用于受控协议测试。"""
from contextlib import contextmanager
import hashlib
import secrets
from sqlalchemy import create_engine, text
SCHEMA = [
"CREATE TABLE IF NOT EXISTS schema_version (version INTEGER NOT NULL)",
"CREATE TABLE IF NOT EXISTS users (id TEXT PRIMARY KEY, username TEXT UNIQUE NOT NULL, password TEXT NOT NULL)",
"CREATE TABLE IF NOT EXISTS devices (id TEXT PRIMARY KEY, user_id TEXT NOT NULL, name TEXT NOT NULL, revoked INTEGER NOT NULL DEFAULT 0)",
"CREATE TABLE IF NOT EXISTS sessions (token TEXT PRIMARY KEY, refresh TEXT UNIQUE NOT NULL, device_id TEXT NOT NULL, expires BIGINT NOT NULL, refresh_expires BIGINT NOT NULL)",
"CREATE TABLE IF NOT EXISTS vaults (id TEXT PRIMARY KEY, user_id TEXT NOT NULL, name TEXT NOT NULL, sequence BIGINT NOT NULL DEFAULT 0, quota BIGINT NOT NULL, used BIGINT NOT NULL DEFAULT 0)",
"CREATE TABLE IF NOT EXISTS uploads (id TEXT PRIMARY KEY, vault_id TEXT NOT NULL, device_id TEXT NOT NULL, hash TEXT NOT NULL, size BIGINT NOT NULL, offset_bytes BIGINT NOT NULL DEFAULT 0, expires BIGINT NOT NULL)",
"CREATE TABLE IF NOT EXISTS objects (vault_id TEXT NOT NULL, hash TEXT NOT NULL, size BIGINT NOT NULL, created BIGINT NOT NULL, PRIMARY KEY(vault_id, hash))",
"CREATE TABLE IF NOT EXISTS revisions (vault_id TEXT NOT NULL, sequence BIGINT NOT NULL, file_id TEXT NOT NULL, base_revision BIGINT NOT NULL, path TEXT NOT NULL, path_key TEXT NOT NULL, operation TEXT NOT NULL, hash TEXT, size BIGINT NOT NULL, device_id TEXT NOT NULL, operation_id TEXT NOT NULL, fingerprint TEXT NOT NULL, PRIMARY KEY(vault_id, sequence), UNIQUE(vault_id, operation_id))",
"CREATE TABLE IF NOT EXISTS files (vault_id TEXT NOT NULL, file_id TEXT NOT NULL, sequence BIGINT NOT NULL, path_key TEXT NOT NULL, deleted INTEGER NOT NULL, PRIMARY KEY(vault_id, file_id))",
"CREATE TABLE IF NOT EXISTS login_limits (key TEXT PRIMARY KEY, started BIGINT NOT NULL, attempts INTEGER NOT NULL)",
]
def password_hash(password: str, salt: str | None = None) -> str:
salt = salt or secrets.token_hex(16)
result = hashlib.scrypt(password.encode(), salt=bytes.fromhex(salt), n=16384, r=8, p=1)
return salt + ":" + result.hex()
class Database:
def __init__(self, url: str):
self.engine = create_engine(url)
self.sqlite = self.engine.dialect.name == "sqlite"
def migrate(self):
with self.transaction() as conn:
conn.execute(text(SCHEMA[0]))
version = conn.execute(text("SELECT version FROM schema_version")).scalar()
if version not in {None, 1}:
raise RuntimeError("数据库版本不兼容,禁止写入")
for statement in SCHEMA[1:]:
conn.execute(text(statement))
if version is None:
conn.execute(text("INSERT INTO schema_version VALUES (1)"))
@contextmanager
def transaction(self):
with self.engine.connect() as conn:
try:
if self.sqlite:
conn.exec_driver_sql("BEGIN IMMEDIATE")
yield conn
conn.commit()
except BaseException:
conn.rollback()
raise
def add_user(self, username: str, password: str):
if len(password) < 12:
raise ValueError("密码至少 12 字符")
with self.transaction() as conn:
conn.execute(text("INSERT INTO users VALUES (:id,:name,:password)"),
{"id": secrets.token_hex(16), "name": username, "password": password_hash(password)})
def row(conn, sql, **params):
return conn.execute(text(sql), params).mappings().first()
def rows(conn, sql, **params):
return conn.execute(text(sql), params).mappings().all()
def run(conn, sql, **params):
return conn.execute(text(sql), params)
+56
View File
@@ -0,0 +1,56 @@
"""协议 v1 DTO;路径在所有平台使用同一套保守规范。"""
import re
import unicodedata
from typing import Literal
from pydantic import BaseModel, ConfigDict, Field, field_validator
class DTO(BaseModel):
model_config = ConfigDict(extra="forbid")
class Login(DTO):
username: str = Field(min_length=1, max_length=80)
password: str = Field(min_length=12, max_length=256)
device_name: str = Field(min_length=1, max_length=120)
class Refresh(DTO):
refresh_token: str = Field(min_length=32, max_length=256)
class VaultCreate(DTO):
name: str = Field(min_length=1, max_length=120)
class Upload(DTO):
content_hash: str = Field(pattern=r"^[a-f0-9]{64}$")
size: int = Field(ge=0, le=104857600)
def canonical_path(value: str) -> str:
if value != unicodedata.normalize("NFC", value) or len(value.encode("utf-8")) > 768:
raise ValueError("路径须为 NFC 且不超过 768 字节")
for part in value.split("/"):
if (not part or part in {".", ".."} or part[-1:] in {" ", "."}
or re.search(r'[<>:"\\|?*\x00-\x1f\x7f]', part)
or re.fullmatch(r"(?i)(CON|PRN|AUX|NUL|COM[1-9]|LPT[1-9])(?:\..*)?", part)
or part.casefold() in {".ainote", ".git"}):
raise ValueError("不可移植或保留路径")
return value
class Commit(DTO):
operation_id: str = Field(pattern=r"^[a-zA-Z0-9-]{16,80}$")
file_id: str = Field(pattern=r"^[a-zA-Z0-9-]{16,80}$")
base_revision: int = Field(ge=0)
path: str = Field(min_length=1)
operation: Literal["put", "delete"]
content_hash: str | None = Field(default=None, pattern=r"^[a-f0-9]{64}$")
size: int = Field(default=0, ge=0, le=104857600)
@field_validator("path")
@classmethod
def path_valid(cls, value):
return canonical_path(value)
+50
View File
@@ -0,0 +1,50 @@
"""内容对象按 Vault 分区;测试磁盘适配器不作为生产对象存储。"""
from pathlib import Path
import hashlib
import os
import tempfile
class DiskObjects:
def __init__(self, root: Path):
self.root = root
root.mkdir(parents=True, exist_ok=True)
def put(self, key: str, data: bytes):
target = self.root / key
target.parent.mkdir(parents=True, exist_ok=True)
fd, temp = tempfile.mkstemp(dir=target.parent)
try:
with os.fdopen(fd, "wb") as stream:
stream.write(data)
stream.flush()
os.fsync(stream.fileno())
os.replace(temp, target)
finally:
Path(temp).unlink(missing_ok=True)
def get(self, key: str) -> bytes:
return (self.root / key).read_bytes()
def delete(self, key: str):
(self.root / key).unlink(missing_ok=True)
class S3Objects:
def __init__(self, endpoint: str, bucket: str):
import boto3
self.client = boto3.client("s3", endpoint_url=endpoint)
self.bucket = bucket
def put(self, key: str, data: bytes):
self.client.put_object(Bucket=self.bucket, Key=key, Body=data,
Metadata={"sha256": hashlib.sha256(data).hexdigest()})
def get(self, key: str) -> bytes:
response = self.client.get_object(Bucket=self.bucket, Key=key)
with response["Body"] as stream:
return stream.read()
def delete(self, key: str):
self.client.delete_object(Bucket=self.bucket, Key=key)