csrf + inor fixes

frdel committed Jul 2, 2025 at 17:03 UTC f51766778b28199d6528156ec4f87ddeafa6ec15
8 files changed +49 -85
models.py
+22 -30
@@ -32,10 +32,16 @@ from langchain_core.messages import (
32 SystemMessage,
33 )
34 from langchain.embeddings.base import Embeddings
35 +from sentence_transformers import SentenceTransformer
36
37 +# disable extra logging
38 +def turn_off_logging():
39 + os.environ['LITELLM_LOG'] = "ERROR" # only errors
40 + litellm.suppress_debug_info = True
41 +
42 +# init
43 load_dotenv()
37 -os.environ['LITELLM_LOG'] = "ERROR" # only errors
38 -litellm.suppress_debug_info = True
44 +turn_off_logging()
45
46
47 class ModelType(Enum):
@@ -125,16 +131,6 @@ def get_rate_limiter(
131 return limiter
132
133
128 -def parse_chunk(chunk: Any):
129 - if isinstance(chunk, str):
130 - content = chunk
131 - elif hasattr(chunk, "content"):
132 - content = str(chunk.content)
133 - else:
134 - content = str(chunk)
135 - return content
136 -
137 -
134 def _parse_chunk(chunk: Any) -> ChatChunk:
135 delta = chunk["choices"][0].get("delta", {})
136 message = chunk["choices"][0].get("model_extra", {}).get("message", {})
@@ -339,6 +335,9 @@ class BrowserCompatibleChatWrapper(LiteLLMChatWrapper):
335 A wrapper for browser agent that can filter/sanitize messages
336 before sending them to the LLM.
337 """
338 + def __init__(self, *args, **kwargs):
339 + turn_off_logging()
340 + super().__init__(*args, **kwargs)
341
342 def _call(
343 self,
@@ -347,7 +346,7 @@ class BrowserCompatibleChatWrapper(LiteLLMChatWrapper):
346 run_manager: Optional[CallbackManagerForLLMRun] = None,
347 **kwargs: Any,
348 ) -> str:
350 - # In the future, message filtering logic can be added here.
349 + turn_off_logging()
350 result = super()._call(messages, stop, run_manager, **kwargs)
351 return result
352
@@ -358,7 +357,7 @@ class BrowserCompatibleChatWrapper(LiteLLMChatWrapper):
357 run_manager: Optional[AsyncCallbackManagerForLLMRun] = None,
358 **kwargs: Any,
359 ) -> AsyncIterator[ChatGenerationChunk]:
361 - # In the future, message filtering logic can be added here.
360 + turn_off_logging()
361 async for chunk in super()._astream(messages, stop, run_manager, **kwargs):
362 yield chunk
363
@@ -374,27 +373,20 @@ class LiteLLMEmbeddingWrapper(Embeddings):
373 def embed_documents(self, texts: List[str]) -> List[List[float]]:
374 resp = embedding(model=self.model_name, input=texts, **self.kwargs)
375 return [
377 - item.get("embedding") if isinstance(item, dict) else item.embedding
378 - for item in resp.data
376 + item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
377 + for item in resp.data # type: ignore
378 ]
379
380 def embed_query(self, text: str) -> List[float]:
381 resp = embedding(model=self.model_name, input=[text], **self.kwargs)
383 - item = resp.data[0]
384 - return item.get("embedding") if isinstance(item, dict) else item.embedding
382 + item = resp.data[0] # type: ignore
383 + return item.get("embedding") if isinstance(item, dict) else item.embedding # type: ignore
384
385
386 class LocalSentenceTransformerWrapper(Embeddings):
387 """Local wrapper for sentence-transformers models to avoid HuggingFace API calls"""
388
389 def __init__(self, model_name: str, **kwargs: Any):
391 - try:
392 - from sentence_transformers import SentenceTransformer
393 - except ImportError:
394 - raise ImportError(
395 - "sentence-transformers library is required for local embeddings. Install with: pip install sentence-transformers"
396 - )
397 -
390 # Remove the "sentence-transformers/" prefix if present
391 if model_name.startswith("sentence-transformers/"):
392 model_name = model_name[len("sentence-transformers/") :]
@@ -403,15 +395,15 @@ class LocalSentenceTransformerWrapper(Embeddings):
395 self.model_name = model_name
396
397 def embed_documents(self, texts: List[str]) -> List[List[float]]:
406 - embeddings = self.model.encode(texts, convert_to_tensor=False)
407 - return embeddings.tolist() if hasattr(embeddings, "tolist") else embeddings
398 + embeddings = self.model.encode(texts, convert_to_tensor=False) # type: ignore
399 + return embeddings.tolist() if hasattr(embeddings, "tolist") else embeddings # type: ignore
400
401 def embed_query(self, text: str) -> List[float]:
410 - embedding = self.model.encode([text], convert_to_tensor=False)
402 + embedding = self.model.encode([text], convert_to_tensor=False) # type: ignore
403 result = (
404 embedding[0].tolist() if hasattr(embedding[0], "tolist") else embedding[0]
405 )
414 - return result
406 + return result # type: ignore
407
408
409 def _get_litellm_chat(
@@ -518,4 +510,4 @@ def _normalize_chat_kwargs(kwargs: Any) -> Any:
510
511
512 def _normalize_embedding_kwargs(kwargs: Any) -> Any:
521 - return kwargs
513 + return kwargs
\ No newline at end of file
prompts/agent0/agent.system.tool.response.md
+1
@@ -16,6 +16,7 @@ usage:
16 "thoughts": [
17 "...",
18 ],
19 + "headline": "Explaining why...",
20 "tool_name": "response",
21 "tool_args": {
22 "text": "Answer to the user",
python/api/csrf_token.py
+10 -3
@@ -1,6 +1,13 @@
1 import secrets
2 -from python.helpers.api import ApiHandler, Input, Output, Request, Response, session
3 -
2 +from python.helpers.api import (
3 + ApiHandler,
4 + Input,
5 + Output,
6 + Request,
7 + Response,
8 + session,
9 +)
10 +from python.helpers import runtime
11
12 class GetCsrfToken(ApiHandler):
13
@@ -15,4 +22,4 @@ class GetCsrfToken(ApiHandler):
22 async def process(self, input: Input, request: Request) -> Output:
23 if "csrf_token" not in session:
24 session["csrf_token"] = secrets.token_urlsafe(32)
18 - return {"token": session["csrf_token"]}
25 + return {"token": session["csrf_token"], "runtime_id": runtime.get_runtime_id()}
python/extensions/response_stream/_10_log_from_stream.py
+3 -4
@@ -21,15 +21,14 @@ class LogFromStream(Extension):
21 heading = build_default_heading(self.agent)
22 if "headline" in parsed:
23 heading = build_heading(self.agent, parsed['headline'])
24 + elif "tool_name" in parsed:
25 + heading = build_heading(self.agent, f"Using tool {parsed['tool_name']}") # if the llm skipped headline
26 elif "thoughts" in parsed:
27 # thought length indicator
28 thoughts = "\n".join(parsed["thoughts"])
29 pipes = "|" * math.ceil(math.sqrt(len(thoughts)))
30 heading = build_heading(self.agent, f"Thinking... {pipes}")
29 -
30 - # if "tool_name" in parsed:
31 - # heading += f" ({parsed['tool_name']})"
32 -
31 +
32 # create log message and store it in loop data temporary params
33 if "log_item_generating" not in loop_data.params_temporary:
34 loop_data.params_temporary["log_item_generating"] = (
python/helpers/runtime.py
+8 -1
@@ -1,5 +1,6 @@
1 import argparse
2 import inspect
3 +import secrets
4 from typing import TypeVar, Callable, Awaitable, Union, overload, cast
5 from python.helpers import dotenv, rfc, settings
6 import asyncio
@@ -12,6 +13,7 @@ R = TypeVar('R')
13 parser = argparse.ArgumentParser()
14 args = {}
15 dockerman = None
16 +runtime_id = None
17
18
19 def initialize():
@@ -38,7 +40,6 @@ def initialize():
40 key = key.lstrip("-")
41 args[key] = value
42
41 -
43 def get_arg(name: str):
44 global args
45 return args.get(name, None)
@@ -58,6 +59,12 @@ def get_local_url():
59 return "host.docker.internal"
60 return "127.0.0.1"
61
62 +def get_runtime_id() -> str:
63 + global runtime_id
64 + if not runtime_id:
65 + runtime_id = secrets.token_hex(8)
66 + return runtime_id
67 +
68 @overload
69 async def call_development_function(func: Callable[..., Awaitable[T]], *args, **kwargs) -> T: ...
70
python/tools/webpage_content_tool._py deleted
-42
@@ -1,42 +0,0 @@
1 -import requests
2 -from bs4 import BeautifulSoup
3 -from urllib.parse import urlparse
4 -from newspaper import Article
5 -from python.helpers.tool import Tool, Response
6 -from python.helpers.errors import handle_error
7 -
8 -
9 -class WebpageContentTool(Tool):
10 - async def execute(self, url="", **kwargs):
11 - if not url:
12 - return Response(message="Error: No URL provided.", break_loop=False)
13 -
14 - try:
15 - # Validate URL
16 - parsed_url = urlparse(url)
17 - if not all([parsed_url.scheme, parsed_url.netloc]):
18 - return Response(message="Error: Invalid URL format.", break_loop=False)
19 -
20 - # Fetch webpage content
21 - response = requests.get(url, timeout=10)
22 - response.raise_for_status()
23 -
24 - # Use newspaper3k for article extraction
25 - article = Article(url)
26 - article.download()
27 - article.parse()
28 -
29 - # If it's not an article, fall back to BeautifulSoup
30 - if not article.text:
31 - soup = BeautifulSoup(response.content, 'html.parser')
32 - text_content = ' '.join(soup.stripped_strings)
33 - else:
34 - text_content = article.text
35 -
36 - return Response(message=f"Webpage content:\n\n{text_content}", break_loop=False)
37 -
38 - except requests.RequestException as e:
39 - return Response(message=f"Error fetching webpage: {str(e)}", break_loop=False)
40 - except Exception as e:
41 - handle_error(e)
42 - return Response(message=f"An error occurred: {str(e)}", break_loop=False)
\ No newline at end of file
run_ui.py
+3 -3
@@ -30,10 +30,10 @@ webapp = Flask("app", static_folder=get_abs_path("./webui"), static_url_path="/"
30 webapp.secret_key = os.getenv("FLASK_SECRET_KEY") or secrets.token_hex(32)
31 webapp.config.update(
32 JSON_SORT_KEYS=False,
33 - SESSION_COOKIE_NAME="session_" + secrets.token_hex(8), # randomize the session cookie name to prevent session collision on same host
33 + SESSION_COOKIE_NAME="session_" + runtime.get_runtime_id(), # bind the session cookie name to runtime id to prevent session collision on same host
34 SESSION_COOKIE_SAMESITE="Strict",
35 SESSION_PERMANENT=True,
36 - PERMANENT_SESSION_LIFETIME=timedelta(days=7)
36 + PERMANENT_SESSION_LIFETIME=timedelta(days=1)
37 )
38
39
@@ -135,7 +135,7 @@ def csrf_protect(f):
135 async def decorated(*args, **kwargs):
136 token = session.get("csrf_token")
137 header = request.headers.get("X-CSRF-Token")
138 - cookie = request.cookies.get("csrf_token")
138 + cookie = request.cookies.get("csrf_token_" + runtime.get_runtime_id())
139 sent = header or cookie
140 if not token or not sent or token != sent:
141 return Response("CSRF token missing or invalid", 403)
webui/js/api.js
+2 -2
@@ -79,6 +79,6 @@ async function getCsrfToken() {
79 credentials: "same-origin",
80 }).then((r) => r.json());
81 csrfToken = response.token;
82 - document.cookie = "csrf_token=" + csrfToken + "; SameSite=Strict; Path=/";
82 + document.cookie = `csrf_token_${response.runtime_id}=${csrfToken}; SameSite=Strict; Path=/`;
83 return csrfToken;
84 -}
\ No newline at end of file
84 +}