Coverage for open_webui/utils/code_interpreter.py: 17%

115 statements  

« prev     ^ index     » next       coverage.py v7.15.2, created at 2026-10-07 05:07 +0000

1import asyncio 

2import logging 

3import uuid 

4from typing import Optional 

5 

6import aiohttp 

7import websockets 

8from open_webui.env import AIOHTTP_CLIENT_ALLOW_REDIRECTS 

9from open_webui.utils.json_codec import JSONCodec 

10from pydantic import BaseModel 

11 

12logger = logging.getLogger(__name__) 

13 

14 

15class ResultModel(BaseModel): 

16 """ 

17 Execute Code Result Model 

18 """ 

19 

20 stdout: Optional[str] = '' 

21 stderr: Optional[str] = '' 

22 result: Optional[str] = '' 

23 

24 

25class JupyterCodeExecuter: 

26 """ 

27 Execute code in jupyter notebook 

28 """ 

29 

30 def __init__( 

31 self, 

32 base_url: str, 

33 code: str, 

34 token: str = '', 

35 password: str = '', 

36 timeout: int = 60, 

37 ): 

38 """ 

39 :param base_url: Jupyter server URL (e.g., "http://localhost:8888") 

40 :param code: Code to execute 

41 :param token: Jupyter authentication token (optional) 

42 :param password: Jupyter password (optional) 

43 :param timeout: WebSocket timeout in seconds (default: 60s) 

44 """ 

45 self.base_url = base_url 

46 self.code = code 

47 self.token = token 

48 self.password = password 

49 self.timeout = timeout 

50 self.kernel_id = '' 

51 if self.base_url[-1] != '/': 

52 self.base_url += '/' 

53 self.session = aiohttp.ClientSession(trust_env=True, base_url=self.base_url) 

54 self.params = {} 

55 self.result = ResultModel() 

56 

57 async def __aenter__(self): 

58 return self 

59 

60 async def __aexit__(self, exc_type, exc_val, exc_tb): 

61 if self.kernel_id: 

62 try: 

63 async with self.session.delete(f'api/kernels/{self.kernel_id}', params=self.params) as response: 

64 response.raise_for_status() 

65 except Exception as err: 

66 logger.exception('close kernel failed, %s', err) 

67 await self.session.close() 

68 

69 async def run(self) -> ResultModel: 

70 try: 

71 await self.sign_in() 

72 await self.init_kernel() 

73 await self.execute_code() 

74 except Exception as err: 

75 logger.exception('execute code failed, %s', err) 

76 self.result.stderr = f'Error: {err}' 

77 return self.result 

78 

79 async def sign_in(self) -> None: 

80 # password authentication 

81 if self.password and not self.token: 

82 async with self.session.get('login') as response: 

83 response.raise_for_status() 

84 xsrf_token = response.cookies['_xsrf'].value 

85 if not xsrf_token: 

86 raise ValueError('_xsrf token not found') 

87 self.session.cookie_jar.update_cookies(response.cookies) 

88 self.session.headers.update({'X-XSRFToken': xsrf_token}) 

89 async with self.session.post( 

90 'login', 

91 data={'_xsrf': xsrf_token, 'password': self.password}, 

92 allow_redirects=AIOHTTP_CLIENT_ALLOW_REDIRECTS, 

93 ) as response: 

94 response.raise_for_status() 

95 self.session.cookie_jar.update_cookies(response.cookies) 

96 

97 # token authentication 

98 if self.token: 

99 self.params.update({'token': self.token}) 

100 

101 async def init_kernel(self) -> None: 

102 async with self.session.post(url='api/kernels', params=self.params) as response: 

103 response.raise_for_status() 

104 kernel_data = await response.json() 

105 self.kernel_id = kernel_data['id'] 

106 

107 def init_ws(self) -> (str, dict): 

108 ws_base = self.base_url.replace('http', 'ws', 1) 

109 ws_params = '?' + '&'.join([f'{key}={val}' for key, val in self.params.items()]) 

110 websocket_url = f'{ws_base}api/kernels/{self.kernel_id}/channels{ws_params if len(ws_params) > 1 else ""}' 

111 ws_headers = {} 

112 if self.password and not self.token: 

113 ws_headers = { 

114 'Cookie': '; '.join([f'{cookie.key}={cookie.value}' for cookie in self.session.cookie_jar]), 

115 **self.session.headers, 

116 } 

117 return websocket_url, ws_headers 

118 

119 async def execute_code(self) -> None: 

120 # initialize ws 

121 websocket_url, ws_headers = self.init_ws() 

122 # execute 

123 async with websockets.connect(websocket_url, additional_headers=ws_headers) as ws: 

124 await self.execute_in_jupyter(ws) 

125 

126 async def execute_in_jupyter(self, ws) -> None: 

127 # send message 

128 msg_id = uuid.uuid4().hex 

129 await ws.send( 

130 JSONCodec.dumps( 

131 { 

132 'header': { 

133 'msg_id': msg_id, 

134 'msg_type': 'execute_request', 

135 'username': 'user', 

136 'session': uuid.uuid4().hex, 

137 'date': '', 

138 'version': '5.3', 

139 }, 

140 'parent_header': {}, 

141 'metadata': {}, 

142 'content': { 

143 'code': self.code, 

144 'silent': False, 

145 'store_history': True, 

146 'user_expressions': {}, 

147 'allow_stdin': False, 

148 'stop_on_error': True, 

149 }, 

150 'channel': 'shell', 

151 } 

152 ) 

153 ) 

154 # parse message 

155 stdout, stderr, result = '', '', [] 

156 while True: 

157 try: 

158 # wait for message 

159 message = await asyncio.wait_for(ws.recv(), self.timeout) 

160 message_data = JSONCodec.loads(message) 

161 # msg id not match, skip 

162 if message_data.get('parent_header', {}).get('msg_id') != msg_id: 

163 continue 

164 # check message type 

165 msg_type = message_data.get('msg_type') 

166 match msg_type: 

167 case 'stream': 

168 if message_data['content']['name'] == 'stdout': 

169 stdout += message_data['content']['text'] 

170 elif message_data['content']['name'] == 'stderr': 

171 stderr += message_data['content']['text'] 

172 case 'execute_result' | 'display_data': 

173 data = message_data['content']['data'] 

174 if 'image/png' in data: 

175 result.append(f'data:image/png;base64,{data["image/png"]}') 

176 elif 'text/plain' in data: 

177 result.append(data['text/plain']) 

178 case 'error': 

179 stderr += '\n'.join(message_data['content']['traceback']) 

180 case 'status': 

181 if message_data['content']['execution_state'] == 'idle': 

182 break 

183 

184 except asyncio.TimeoutError: 

185 stderr += '\nExecution timed out.' 

186 break 

187 self.result.stdout = stdout.strip() 

188 self.result.stderr = stderr.strip() 

189 self.result.result = '\n'.join(result).strip() if result else '' 

190 

191 

192async def execute_code_jupyter( 

193 base_url: str, code: str, token: str = '', password: str = '', timeout: int = 60 

194) -> dict: 

195 async with JupyterCodeExecuter(base_url, code, token, password, timeout) as executor: 

196 result = await executor.run() 

197 return result.model_dump()