token_store.py 6.4 KB

123456789101112131415161718192021222324252627282930313233343536373839404142434445464748495051525354555657585960616263646566676869707172737475767778798081828384858687888990919293949596979899100101102103104105106107108109110111112113114115116117118119120121122123124125126127128129130131132133134135136137138139140141142143144145146147148149150151152153154155156157158159160161162163164165166167168169170171172173174175176177178179180181182183184185186187188189190191
  1. from datetime import datetime, timedelta
  2. import json
  3. import os
  4. import socket
  5. class InMemoryTokenStore:
  6. def __init__(self, refresh_skew_seconds=120):
  7. self.refresh_skew_seconds = int(refresh_skew_seconds)
  8. self._session = None
  9. def save(self, token, expire_time):
  10. self._session = {
  11. 'token': token,
  12. 'expire_time': expire_time,
  13. }
  14. return self._session
  15. def get(self):
  16. return self._session
  17. def clear(self):
  18. self._session = None
  19. def is_expiring(self):
  20. if not self._session:
  21. return True
  22. expire_time = self._session.get('expire_time')
  23. if not expire_time:
  24. return True
  25. if isinstance(expire_time, str) and expire_time.endswith('Z'):
  26. expire_time = expire_time[:-1] + '+00:00'
  27. expires_at = datetime.fromisoformat(str(expire_time))
  28. return expires_at <= datetime.now(expires_at.tzinfo) + timedelta(seconds=self.refresh_skew_seconds)
  29. def require_token(self):
  30. if not self._session or not self._session.get('token'):
  31. raise RuntimeError('mcp token missing')
  32. return self._session['token']
  33. class FileTokenStore(InMemoryTokenStore):
  34. def __init__(self, path, refresh_skew_seconds=120):
  35. super().__init__(refresh_skew_seconds=refresh_skew_seconds)
  36. self.path = path
  37. self._session = self._read_session()
  38. def save(self, token, expire_time):
  39. session = super().save(token, expire_time)
  40. directory = os.path.dirname(self.path)
  41. if directory:
  42. os.makedirs(directory, exist_ok=True)
  43. with open(self.path, 'w', encoding='utf-8') as handle:
  44. json.dump(session, handle, ensure_ascii=False)
  45. return session
  46. def get(self):
  47. if self._session is None and os.path.exists(self.path):
  48. self._session = self._read_session()
  49. return self._session
  50. def clear(self):
  51. super().clear()
  52. if os.path.exists(self.path):
  53. os.remove(self.path)
  54. def _read_session(self):
  55. if not os.path.exists(self.path):
  56. return None
  57. with open(self.path, 'r', encoding='utf-8') as handle:
  58. return json.load(handle)
  59. class RedisTokenStore(InMemoryTokenStore):
  60. def __init__(self, client, prefix='fms:mcp:workbuddy:', session_key='', refresh_skew_seconds=120):
  61. super().__init__(refresh_skew_seconds=refresh_skew_seconds)
  62. session_key = str(session_key or '').strip()
  63. if not session_key:
  64. raise ValueError('session_key is required for RedisTokenStore')
  65. self.client = client
  66. self.prefix = str(prefix or 'fms:mcp:workbuddy:')
  67. self.session_key = session_key
  68. self.key = self.prefix + self.session_key
  69. def save(self, token, expire_time):
  70. session = super().save(token, expire_time)
  71. payload = json.dumps(session, ensure_ascii=False)
  72. ttl = self._ttl_seconds(expire_time)
  73. self.client.set(self.key, payload, ex=ttl if ttl and ttl > 0 else None)
  74. return session
  75. def get(self):
  76. raw = self.client.get(self.key)
  77. if raw is None or raw == '':
  78. self._session = None
  79. return None
  80. if isinstance(raw, bytes):
  81. raw = raw.decode('utf-8')
  82. self._session = json.loads(raw)
  83. return self._session
  84. def clear(self):
  85. super().clear()
  86. self.client.delete(self.key)
  87. def is_expiring(self):
  88. self.get()
  89. return super().is_expiring()
  90. @staticmethod
  91. def _ttl_seconds(expire_time):
  92. if not expire_time:
  93. return None
  94. value = expire_time[:-1] + '+00:00' if isinstance(expire_time, str) and expire_time.endswith('Z') else expire_time
  95. expires_at = datetime.fromisoformat(str(value))
  96. now = datetime.now(expires_at.tzinfo) if expires_at.tzinfo else datetime.now()
  97. return max(1, int((expires_at - now).total_seconds()))
  98. class RedisSocketClient:
  99. def __init__(self, host='127.0.0.1', port=6379, db=0, password='', timeout=5):
  100. self.host = host
  101. self.port = int(port)
  102. self.db = int(db)
  103. self.password = password or ''
  104. self.timeout = int(timeout)
  105. def set(self, key, value, ex=None):
  106. command = ['SET', key, value]
  107. if ex is not None:
  108. command.extend(['EX', int(ex)])
  109. return self._execute(*command) == 'OK'
  110. def get(self, key):
  111. return self._execute('GET', key)
  112. def delete(self, key):
  113. return self._execute('DEL', key)
  114. def _execute(self, *parts):
  115. with socket.create_connection((self.host, self.port), timeout=self.timeout) as sock:
  116. stream = sock.makefile('rb')
  117. if self.password:
  118. self._send(sock, 'AUTH', self.password)
  119. self._read_response(stream)
  120. if self.db:
  121. self._send(sock, 'SELECT', self.db)
  122. self._read_response(stream)
  123. self._send(sock, *parts)
  124. return self._read_response(stream)
  125. @staticmethod
  126. def _send(sock, *parts):
  127. payload = ['*{0}\r\n'.format(len(parts)).encode('utf-8')]
  128. for part in parts:
  129. if isinstance(part, bytes):
  130. data = part
  131. else:
  132. data = str(part).encode('utf-8')
  133. payload.append('${0}\r\n'.format(len(data)).encode('utf-8'))
  134. payload.append(data + b'\r\n')
  135. sock.sendall(b''.join(payload))
  136. def _read_response(self, stream):
  137. prefix = stream.read(1)
  138. if not prefix:
  139. raise RuntimeError('empty redis response')
  140. if prefix == b'+':
  141. return self._read_line(stream).decode('utf-8')
  142. if prefix == b'-':
  143. raise RuntimeError(self._read_line(stream).decode('utf-8'))
  144. if prefix == b':':
  145. return int(self._read_line(stream))
  146. if prefix == b'$':
  147. length = int(self._read_line(stream))
  148. if length == -1:
  149. return None
  150. data = stream.read(length)
  151. stream.read(2)
  152. return data.decode('utf-8')
  153. if prefix == b'*':
  154. length = int(self._read_line(stream))
  155. return [self._read_response(stream) for _ in range(length)]
  156. raise RuntimeError('unsupported redis response: {0}'.format(prefix))
  157. @staticmethod
  158. def _read_line(stream):
  159. line = stream.readline()
  160. if not line.endswith(b'\r\n'):
  161. raise RuntimeError('invalid redis line')
  162. return line[:-2]