Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
27 changes: 20 additions & 7 deletions lark_channel/card/action_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -2,7 +2,7 @@
import hmac
import json
import logging
from typing import Optional, Callable, Any, TYPE_CHECKING
from typing import Optional, Callable, Any, Dict, TYPE_CHECKING

from lark_channel.core.const import *
from lark_channel.core.enum import LogLevel
Expand Down Expand Up @@ -179,9 +179,9 @@ def _preverify_encrypted_request(self, request: RawRequest) -> bool:

def _has_signature_headers(self, request: RawRequest) -> bool:
return (
Strings.is_not_empty(request.headers.get(LARK_REQUEST_TIMESTAMP))
and Strings.is_not_empty(request.headers.get(LARK_REQUEST_NONCE))
and Strings.is_not_empty(request.headers.get(LARK_REQUEST_SIGNATURE))
Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_TIMESTAMP))
and Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_NONCE))
and Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_SIGNATURE))
)

def _record_security_audit(
Expand All @@ -206,9 +206,9 @@ def _record_security_audit(
def _verify_sign(self, request: RawRequest) -> None:
if self._verification_token is None or self._verification_token == "":
return
timestamp = request.headers.get(LARK_REQUEST_TIMESTAMP)
nonce = request.headers.get(LARK_REQUEST_NONCE)
signature = request.headers.get(LARK_REQUEST_SIGNATURE)
timestamp = _get_header(request.headers, LARK_REQUEST_TIMESTAMP)
nonce = _get_header(request.headers, LARK_REQUEST_NONCE)
signature = _get_header(request.headers, LARK_REQUEST_SIGNATURE)
bs = (timestamp + nonce + self._verification_token).encode(UTF_8) + request.body
h = hashlib.sha1(bs)
if signature != h.hexdigest():
Expand Down Expand Up @@ -260,3 +260,16 @@ def _default_security_config():
from lark_channel.channel.config import SecurityConfig

return SecurityConfig()


def _get_header(headers: Dict[str, str], name: str) -> Optional[str]:
# ASGI servers (Starlette/FastAPI) hand handlers lowercase header names,
# so an exact-case lookup silently misses X-Lark-* signature headers.
value = headers.get(name)
if value is not None:
return value
lname = name.lower()
for key, val in headers.items():
if key.lower() == lname:
return val
return None
113 changes: 113 additions & 0 deletions lark_channel/channel/tests/test_handle_webhook_request.py
Original file line number Diff line number Diff line change
Expand Up @@ -223,6 +223,33 @@ def test_card_callback_with_signature_header_does_not_require_verification_token
assert len(seen) == 1


def test_signed_card_accepts_lowercase_asgi_headers():
"""ASGI servers (Starlette/FastAPI) hand handlers lowercase header
names; a legitimately signed plaintext card callback must not 500."""
seen = []
body = json.dumps(
{
"type": "card.action.trigger",
"action": {"value": {"key": "value"}},
}
).encode("utf-8")
handler = (
CardActionHandler.builder("", "verification-token")
.register(lambda card: seen.append(card))
.build()
)
headers = {
key.lower(): value
for key, value in _signed_headers(body, "verification-token", algorithm="sha1").items()
}

resp = handler.do(_request_bytes(body, headers))

assert resp.status_code == 200
assert resp.content == b'{"msg":"success"}'
assert len(seen) == 1


def test_signed_encrypted_event_is_verified_before_dispatch():
seen = []
body = _encrypted_body(
Expand Down Expand Up @@ -251,6 +278,66 @@ def test_signed_encrypted_event_is_verified_before_dispatch():
assert len(seen) == 1


def test_signed_event_accepts_lowercase_asgi_headers():
"""ASGI servers (Starlette/FastAPI) hand handlers lowercase header
names; a legitimately signed request must not 500 on that alone."""
seen = []
body = json.dumps(
{
"schema": "2.0",
"header": {
"event_type": "example.event",
"token": "verification-token",
},
"event": {"value": "ok"},
}
).encode("utf-8")
handler = (
EventDispatcherHandler.builder("encrypt-key", "verification-token")
.register_p2_customized_event("example.event", lambda event: seen.append(event))
.build()
)
headers = {
key.lower(): value for key, value in _signed_headers(body, "encrypt-key").items()
}

resp = handler.do(_request_bytes(body, headers))

assert resp.status_code == 200
assert resp.content == b'{"msg":"success"}'
assert len(seen) == 1


def test_signed_encrypted_event_accepts_lowercase_asgi_headers():
seen = []
body = _encrypted_body(
{
"schema": "2.0",
"header": {
"event_type": "example.event",
"token": "verification-token",
},
"event": {"value": "ok"},
},
"encrypt-key",
)
handler = (
EventDispatcherHandler.builder("encrypt-key", "verification-token")
.register_p2_customized_event("example.event", lambda event: seen.append(event))
.build()
)
headers = {
key.lower(): value
for key, value in _signed_headers(body, "encrypt-key", algorithm="sha256").items()
}

resp = handler.do(_request_bytes(body, headers))

assert resp.status_code == 200
assert resp.content == b'{"msg":"success"}'
assert len(seen) == 1


def test_strict_event_invalid_signature_rejects_before_decrypt(monkeypatch):
recorder = InMemorySecurityAuditRecorder()
body = _encrypted_body({"type": "url_verification"}, "encrypt-key")
Expand Down Expand Up @@ -441,6 +528,32 @@ def test_signed_encrypted_card_is_verified_before_dispatch():
assert len(seen) == 1


def test_signed_encrypted_card_accepts_lowercase_asgi_headers():
seen = []
body = _encrypted_body(
{
"type": "card.action.trigger",
"action": {"value": {"key": "value"}},
},
"encrypt-key",
)
handler = (
CardActionHandler.builder("encrypt-key", "verification-token")
.register(lambda card: seen.append(card))
.build()
)
headers = {
key.lower(): value
for key, value in _signed_headers(body, "verification-token", algorithm="sha1").items()
}

resp = handler.do(_request_bytes(body, headers))

assert resp.status_code == 200
assert resp.content == b'{"msg":"success"}'
assert len(seen) == 1


def test_strict_card_invalid_signature_rejects_before_decrypt(monkeypatch):
recorder = InMemorySecurityAuditRecorder()
body = _encrypted_body({"type": "card.action.trigger"}, "encrypt-key")
Expand Down
25 changes: 19 additions & 6 deletions lark_channel/event/dispatcher_handler.py
Original file line number Diff line number Diff line change
Expand Up @@ -218,9 +218,9 @@ def _preverify_encrypted_request(self, request: RawRequest) -> bool:

def _has_signature_headers(self, request: RawRequest) -> bool:
return (
Strings.is_not_empty(request.headers.get(LARK_REQUEST_TIMESTAMP))
and Strings.is_not_empty(request.headers.get(LARK_REQUEST_NONCE))
and Strings.is_not_empty(request.headers.get(LARK_REQUEST_SIGNATURE))
Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_TIMESTAMP))
and Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_NONCE))
and Strings.is_not_empty(_get_header(request.headers, LARK_REQUEST_SIGNATURE))
)

def _record_security_audit(
Expand All @@ -245,9 +245,9 @@ def _record_security_audit(
def _verify_sign(self, request: RawRequest) -> None:
if self._encrypt_key is None or self._encrypt_key == "":
return
timestamp = request.headers.get(LARK_REQUEST_TIMESTAMP)
nonce = request.headers.get(LARK_REQUEST_NONCE)
signature = request.headers.get(LARK_REQUEST_SIGNATURE)
timestamp = _get_header(request.headers, LARK_REQUEST_TIMESTAMP)
nonce = _get_header(request.headers, LARK_REQUEST_NONCE)
signature = _get_header(request.headers, LARK_REQUEST_SIGNATURE)
bs = (timestamp + nonce + self._encrypt_key).encode(UTF_8) + request.body
if signature != hashlib.sha256(bs).hexdigest():
raise AccessDeniedException("signature verification failed")
Expand Down Expand Up @@ -423,3 +423,16 @@ def _default_security_config():
from lark_channel.channel.config import SecurityConfig

return SecurityConfig()


def _get_header(headers: Dict[str, str], name: str) -> Optional[str]:
# ASGI servers (Starlette/FastAPI) hand handlers lowercase header names,
# so an exact-case lookup silently misses X-Lark-* signature headers.
value = headers.get(name)
if value is not None:
return value
lname = name.lower()
for key, val in headers.items():
if key.lower() == lname:
return val
return None