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
+}