Skip to content

Commit 84f106f

Browse files
committed
refactor(manager, xray, crypto): improve state serialization and type handling
- Moved state serialization in CoreManager to ensure proper encoding. - Updated XRayConfig to ensure SNI values are converted to strings. - Enhanced get_cert_SANs function to return SAN values as strings. - Added unit tests for get_cert_SANs to verify JSON serializability of returned values.
1 parent 25fe11c commit 84f106f

4 files changed

Lines changed: 57 additions & 10 deletions

File tree

app/core/manager.py

Lines changed: 2 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -61,9 +61,8 @@ async def _persist_state(self):
6161
if not self._kv:
6262
return
6363
state = await self._snapshot_state()
64-
# State is already serialized by _snapshot_state; just encode
65-
state_bytes = json.dumps(state).encode("utf-8")
6664
try:
65+
state_bytes = json.dumps(state).encode("utf-8")
6766
await self._kv.put(self.STATE_CACHE_KEY, state_bytes)
6867
except Exception as exc:
6968
self._logger.warning(f"Failed to persist core state to NATS KV: {exc}")
@@ -80,7 +79,7 @@ async def _load_state_from_cache(self) -> bool:
8079
# Deserialize state using JSON
8180
try:
8281
cached_state = json.loads(entry.value.decode("utf-8"))
83-
except json.JSONDecodeError, UnicodeDecodeError:
82+
except (json.JSONDecodeError, UnicodeDecodeError):
8483
self._logger.warning("Failed to decode CoreManager state as JSON, ignoring...")
8584
return False
8685

app/core/xray.py

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -181,7 +181,7 @@ def _handle_tls_settings(self, tls_settings: dict, settings: dict, inbound_tag:
181181
settings["pinnedPeerCertSha256"] = tls_settings.get("pinnedPeerCertSha256", "")
182182
settings["fp"] = tls_settings.get("fingerprint", "chrome")
183183
if sni := tls_settings.get("serverName"):
184-
settings["sni"].append(sni)
184+
settings["sni"].append(str(sni))
185185
for certificate in tls_settings.get("certificates", []):
186186
serve_on_node = certificate.pop("serveOnNode", False)
187187
if serve_on_node:
@@ -218,7 +218,7 @@ def _handle_reality_settings(self, tls_settings: dict, settings: dict, inbound_t
218218
"""Handle Reality security settings."""
219219
settings["fp"] = tls_settings.get("fingerprint", "chrome")
220220
settings["tls"] = "reality"
221-
settings["sni"] = tls_settings.get("serverNames", [])
221+
settings["sni"] = [str(sni) for sni in (tls_settings.get("serverNames") or [])]
222222

223223
pvk = tls_settings.get("privateKey")
224224
if not pvk:

app/utils/crypto.py

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -9,14 +9,14 @@
99
from cryptography.hazmat.primitives.asymmetric import x25519
1010

1111

12-
def get_cert_SANs(cert: bytes):
12+
def get_cert_SANs(cert: bytes) -> list[str]:
13+
"""Return SAN values as strings (IP SANs are ipaddress objects)."""
1314
cert = x509.load_pem_x509_certificate(cert, default_backend())
14-
san_list = []
15+
san_list: list[str] = []
1516
for extension in cert.extensions:
1617
if isinstance(extension.value, x509.SubjectAlternativeName):
17-
san = extension.value
18-
for name in san:
19-
san_list.append(name.value)
18+
for name in extension.value:
19+
san_list.append(str(name.value))
2020
return san_list
2121

2222

tests/test_cert_sans.py

Lines changed: 48 additions & 0 deletions
Original file line numberDiff line numberDiff line change
@@ -0,0 +1,48 @@
1+
import datetime
2+
import json
3+
from ipaddress import ip_address
4+
5+
from cryptography import x509
6+
from cryptography.hazmat.primitives import hashes, serialization
7+
from cryptography.hazmat.primitives.asymmetric import rsa
8+
from cryptography.x509.oid import NameOID
9+
10+
from app.utils.crypto import get_cert_SANs
11+
12+
13+
def _pem_cert_with_dns_and_ip_sans() -> bytes:
14+
key = rsa.generate_private_key(public_exponent=65537, key_size=2048)
15+
name = x509.Name([x509.NameAttribute(NameOID.COMMON_NAME, "example.com")])
16+
now = datetime.datetime.now(datetime.UTC)
17+
cert = (
18+
x509.CertificateBuilder()
19+
.subject_name(name)
20+
.issuer_name(name)
21+
.public_key(key.public_key())
22+
.serial_number(x509.random_serial_number())
23+
.not_valid_before(now - datetime.timedelta(days=1))
24+
.not_valid_after(now + datetime.timedelta(days=90))
25+
.add_extension(
26+
x509.SubjectAlternativeName(
27+
[
28+
x509.DNSName("example.com"),
29+
x509.DNSName("www.example.com"),
30+
x509.IPAddress(ip_address("1.2.3.4")),
31+
x509.IPAddress(ip_address("2001:db8::1")),
32+
]
33+
),
34+
critical=False,
35+
)
36+
.sign(key, hashes.SHA256())
37+
)
38+
return cert.public_bytes(serialization.Encoding.PEM)
39+
40+
41+
def test_get_cert_sans_returns_json_serializable_strings() -> None:
42+
sans = get_cert_SANs(_pem_cert_with_dns_and_ip_sans())
43+
assert all(isinstance(item, str) for item in sans)
44+
assert "example.com" in sans
45+
assert "www.example.com" in sans
46+
assert "1.2.3.4" in sans
47+
assert "2001:db8::1" in sans
48+
json.dumps({"sni": sans})

0 commit comments

Comments
 (0)