Skip to content

Commit 6fae22b

Browse files
committed
fix: use execute to consume keys used to publish
1 parent 5c6a580 commit 6fae22b

3 files changed

Lines changed: 48 additions & 29 deletions

File tree

grpc_services/v3/utils.py

Lines changed: 23 additions & 25 deletions
Original file line numberDiff line numberDiff line change
@@ -12,7 +12,7 @@
1212
X25519PublicKey,
1313
)
1414
from sqlalchemy import delete, insert, select
15-
from sqlalchemy.orm import Session, contains_eager
15+
from sqlalchemy.orm import Session
1616

1717
from db import get_session
1818
from lib_relaysms_payload_specs.generated import relaysms_spec_payload as rrs
@@ -36,35 +36,33 @@ def get_keys_for_decryption(
3636
Fetch token and all necessary keys for decryption.
3737
Returns (Token, TokenHash, ss_kid, es_kid, es_kid_pk, ec_kid_pk).
3838
"""
39-
stmt = (
40-
select(Token)
41-
.join(Token.token_hash)
42-
.join(TokenHash.server_keys)
43-
.join(TokenHash.client_keys)
44-
.where(Token.token_id == token_id_bytes)
45-
.where(ServerEphemeralKey.key_index == key_id)
46-
.where(ClientEphemeralKey.key_index == key_id)
47-
.options(
48-
contains_eager(Token.token_hash).contains_eager(TokenHash.server_keys),
49-
contains_eager(Token.token_hash).contains_eager(TokenHash.client_keys),
50-
)
51-
)
52-
53-
token = session.scalars(stmt).first()
39+
token = session.scalar(select(Token).where(Token.token_id == token_id_bytes))
5440
if not token:
55-
raise ValueError("Token or associated keys not found")
41+
raise ValueError("token not found")
5642

57-
token_hash_obj = token.token_hash
43+
token_hash_obj = session.scalar(
44+
select(TokenHash).where(TokenHash.id == token.token_hash_id)
45+
)
5846
if not token_hash_obj:
59-
raise ValueError("Token hash not found")
47+
raise ValueError("token hash not found")
6048

61-
if not token_hash_obj.server_keys:
62-
raise ValueError(f"Server ephemeral key not found for kid {key_id}")
63-
se_key = token_hash_obj.server_keys[0]
49+
se_key = session.scalar(
50+
select(ServerEphemeralKey).where(
51+
ServerEphemeralKey.token_hash_id == token_hash_obj.id,
52+
ServerEphemeralKey.key_index == key_id,
53+
)
54+
)
55+
if not se_key:
56+
raise ValueError(f"server ephemeral key not found: kid={key_id}")
6457

65-
if not token_hash_obj.client_keys:
66-
raise ValueError(f"Client ephemeral key not found for kid {key_id}")
67-
ce_key = token_hash_obj.client_keys[0]
58+
ce_key = session.scalar(
59+
select(ClientEphemeralKey).where(
60+
ClientEphemeralKey.token_hash_id == token_hash_obj.id,
61+
ClientEphemeralKey.key_index == key_id,
62+
)
63+
)
64+
if not ce_key:
65+
raise ValueError(f"client ephemeral key not found: kid={key_id}")
6866

6967
ss_kid = get_private_key(key_id, session).private_bytes_raw()
7068

rest_services/v1/routes.py

Lines changed: 10 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -188,9 +188,16 @@ def create_publications(
188188
status_code=400, detail="Failed to deserialize payload"
189189
) from exc
190190

191+
t_id = payload.get_t_id()
192+
if t_id is None:
193+
logger.error("payload is missing token ID")
194+
raise HTTPException(
195+
status_code=400, detail="Payload is missing token ID"
196+
)
197+
191198
with get_session() as db:
192199
publish_content(
193-
token_id=struct.pack("<I", payload.get_t_id()),
200+
token_id=struct.pack("<I", t_id),
194201
key_id=payload.get_kid(),
195202
len_att=payload.get_len_att(),
196203
content_ciphertext=payload.get_content(),
@@ -210,6 +217,7 @@ def create_publications(
210217
db=db,
211218
)
212219
if joined is None:
220+
logger.info("Segment stored. Waiting for remaining segments.")
213221
return PublishContentResponse(
214222
message="Segment stored. Waiting for remaining segments."
215223
)
@@ -229,4 +237,5 @@ def create_publications(
229237
detail=f"Payload type {payload_type!r} is not supported.",
230238
)
231239

240+
logger.info("Content published successfully.")
232241
return PublishContentResponse(message="Content published successfully")

rest_services/v1/services.py

Lines changed: 15 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,6 +1,7 @@
11
# SPDX-License-Identifier: GPL-3.0-only
22
"""v1 Services for the REST API."""
33

4+
from sqlalchemy import delete
45
from sqlalchemy.orm import Session
56

67
from grpc_services.v3.utils import (
@@ -10,11 +11,13 @@
1011
)
1112
from lib_relaysms_payload_specs.generated import relaysms_spec_payload as rrs
1213
from logutils import get_logger
14+
from models.client_ephemeral_key import ClientEphemeralKey
1315
from models.payload_segment import create_if_not_exists as create_segment
1416
from models.payload_segment import get_all_data
1517
from models.payload_session import create as create_session
1618
from models.payload_session import delete as delete_session
1719
from models.payload_session import get_by_sender_and_session
20+
from models.server_ephemeral_key import ServerEphemeralKey
1821
from models.server_identity_key import mark_key_used as mark_ss_kid_used
1922
from models.token import update_token_data
2023
from models.token_hash import TokenHash
@@ -28,7 +31,6 @@ def _get_adapter_params(
2831
token_data: dict, content: rrs.V1ContentsContainer, *, extras: dict | None = None
2932
) -> dict:
3033
cat_id = content.get_cat_id()
31-
print(">>>>>>>>> ATTACHMENT:", content.get_attachment())
3234
extra_params = extras or {}
3335

3436
match cat_id:
@@ -54,8 +56,18 @@ def _get_adapter_params(
5456

5557

5658
def _consume_used_keys(token_hash: TokenHash, key_index: int, session: Session) -> None:
57-
session.delete(token_hash.server_keys[0])
58-
session.delete(token_hash.client_keys[0])
59+
session.execute(
60+
delete(ServerEphemeralKey).where(
61+
ServerEphemeralKey.token_hash_id == token_hash.id,
62+
ServerEphemeralKey.key_index == key_index,
63+
)
64+
)
65+
session.execute(
66+
delete(ClientEphemeralKey).where(
67+
ClientEphemeralKey.token_hash_id == token_hash.id,
68+
ClientEphemeralKey.key_index == key_index,
69+
)
70+
)
5971
mark_ss_kid_used(key_index, session)
6072

6173

0 commit comments

Comments
 (0)