diff --git a/bemani/data/mysql/machine.py b/bemani/data/mysql/machine.py index a51c060..199ac20 100644 --- a/bemani/data/mysql/machine.py +++ b/bemani/data/mysql/machine.py @@ -140,7 +140,7 @@ class MachineData(BaseData): def from_session(self, session: str) -> Optional[ArcadeID]: """ - Given a previously-opened session, look up a user ID. + Given a previously-opened session, look up an arcade ID. Parameters: session - String identifying a session that was opened by create_session. @@ -151,6 +151,10 @@ class MachineData(BaseData): arcadeid = self._from_session(session, "arcadeid") if arcadeid is None: return None + sql = "SELECT id FROM arcade WHERE id = :arcadeid LIMIT 1" + cursor = self.execute(sql, {"arcadeid": arcadeid}) + if cursor.rowcount != 1: + return None return ArcadeID(arcadeid) def get_machine(self, pcbid: str) -> Optional[Machine]: @@ -570,7 +574,7 @@ class MachineData(BaseData): def create_session(self, arcadeid: ArcadeID, expiration: int = (30 * 86400)) -> str: """ - Given a user ID, create a session string. + Given an arcade ID, create a session string. Parameters: arcadeid - Arcade ID we wish to start a session for. diff --git a/bemani/data/mysql/user.py b/bemani/data/mysql/user.py index a2d956d..7a0943f 100644 --- a/bemani/data/mysql/user.py +++ b/bemani/data/mysql/user.py @@ -278,6 +278,11 @@ class UserData(BaseData): userid = self._from_session(session, "userid") if userid is None: return None + sql = "SELECT id FROM user WHERE id = :userid LIMIT 1" + cursor = self.execute(sql, {"userid": userid}) + if cursor.rowcount != 1: + return None + return UserID(userid) def get_user(self, userid: UserID) -> Optional[User]: