55#
66from dataclasses import dataclass , field
77from datetime import datetime , timedelta , timezone
8+ from time import time
89from typing import TypedDict
910
11+ from agora_token_builder import RtcTokenBuilder
1012from ten_runtime import AsyncTenEnv
1113from ten_ai_base .config import BaseConfig
1214from spatius import new_avatar_session , AgoraEgressConfig
@@ -23,6 +25,7 @@ class SpatiusParams(TypedDict, total=False):
2325 agora_uid : str
2426 agora_token : str
2527 agora_appid : str
28+ agora_appcert : str
2629 agora_channel : str
2730 region : str
2831 sample_rate : int | str
@@ -40,12 +43,14 @@ class SpatiusConfig(BaseConfig):
4043 agora_uid : str = ""
4144 agora_token : str = ""
4245 agora_appid : str = ""
46+ agora_appcert : str = ""
4347 agora_channel : str = ""
4448
4549 region : str = ""
4650 sample_rate : int = 24000
4751 session_expire_minutes : int = 30
4852
53+ channel : str = ""
4954 params : SpatiusParams = field (default_factory = dict )
5055
5156 dump : bool = False
@@ -71,9 +76,15 @@ def update_params(self) -> None:
7176 if "agora_appid" in self .params :
7277 self .agora_appid = self .params ["agora_appid" ]
7378
79+ if "agora_appcert" in self .params :
80+ self .agora_appcert = self .params ["agora_appcert" ]
81+
7482 if "agora_channel" in self .params :
7583 self .agora_channel = self .params ["agora_channel" ]
7684
85+ if self ._has_value (self .channel ):
86+ self .agora_channel = self .channel
87+
7788 if "region" in self .params :
7889 self .region = self .params ["region" ]
7990
@@ -92,7 +103,6 @@ def validate_params(self) -> None:
92103 "params.spatius_app_id" : self .spatius_app_id ,
93104 "params.spatius_avatar_id" : self .spatius_avatar_id ,
94105 "params.agora_uid" : self .agora_uid ,
95- "params.agora_token" : self .agora_token ,
96106 "params.agora_appid" : self .agora_appid ,
97107 "params.agora_channel" : self .agora_channel ,
98108 }
@@ -107,6 +117,14 @@ def validate_params(self) -> None:
107117 f"Missing required fields: { ', ' .join (missing_fields )} "
108118 )
109119
120+ if not self ._has_value (self .agora_token ) and not self ._has_value (
121+ self .agora_appcert
122+ ):
123+ raise ValueError (
124+ "Either params.agora_token or params.agora_appcert "
125+ "must be provided"
126+ )
127+
110128 if self .sample_rate <= 0 :
111129 raise ValueError ("sample_rate must be greater than 0" )
112130
@@ -118,6 +136,25 @@ def validate_params(self) -> None:
118136 except ValueError as exc :
119137 raise ValueError ("params.agora_uid must be an integer" ) from exc
120138
139+ def resolve_agora_token (self ) -> str :
140+ """Return configured Agora token or generate one from app cert."""
141+ if self ._has_value (self .agora_token ):
142+ return self .agora_token
143+
144+ privilege_expired_ts = int (time ()) + (self .session_expire_minutes * 60 )
145+ return RtcTokenBuilder .buildTokenWithUid (
146+ self .agora_appid ,
147+ self .agora_appcert ,
148+ self .agora_channel ,
149+ int (self .agora_uid ),
150+ 1 ,
151+ privilege_expired_ts ,
152+ )
153+
154+ @staticmethod
155+ def _has_value (value : str ) -> bool :
156+ return bool (value and value .strip ())
157+
121158
122159class SpatiusAvatarExtension (AsyncAvatarBaseExtension ):
123160 """
@@ -174,6 +211,7 @@ async def validate_config(self, ten_env: AsyncTenEnv) -> bool:
174211 f"agora_uid={ self .config .agora_uid } , "
175212 f"agora_token={ self ._masked_agora_token ()} , "
176213 f"agora_appid={ self .config .agora_appid } , "
214+ f"agora_appcert={ self ._masked_agora_appcert ()} , "
177215 f"agora_channel={ self .config .agora_channel } , "
178216 f"sample_rate={ self .config .sample_rate } , "
179217 "session_expire_minutes="
@@ -193,10 +231,20 @@ def _masked_api_key(self) -> str:
193231
194232 def _masked_agora_token (self ) -> str :
195233 """Return a redacted Agora token for logs."""
234+ if not self .config .agora_token :
235+ return "(generated from app cert)"
196236 if len (self .config .agora_token ) <= 4 :
197237 return "(short)"
198238 return f"***{ self .config .agora_token [- 4 :]} "
199239
240+ def _masked_agora_appcert (self ) -> str :
241+ """Return a redacted Agora app certificate for logs."""
242+ if not self .config .agora_appcert :
243+ return "(empty)"
244+ if len (self .config .agora_appcert ) <= 4 :
245+ return "(short)"
246+ return f"***{ self .config .agora_appcert [- 4 :]} "
247+
200248 def _region (self ) -> str :
201249 """Return the configured Spatius region."""
202250 return (self .config .region or "" ).strip ()
@@ -213,9 +261,10 @@ async def connect_to_avatar(self, ten_env: AsyncTenEnv) -> None:
213261
214262 # Create avatar session using spatius with Agora egress.
215263 avatar_uid = int (self .config .agora_uid )
264+ agora_token = self .config .resolve_agora_token ()
216265 agora_egress = AgoraEgressConfig (
217266 channel_name = self .config .agora_channel ,
218- token = self . config . agora_token ,
267+ token = agora_token ,
219268 uid = avatar_uid ,
220269 publisher_id = self .config .agora_uid ,
221270 )
0 commit comments