Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
78 changes: 64 additions & 14 deletions browser-use-demo/browser_use_demo/loop.py
Original file line number Diff line number Diff line change
Expand Up @@ -6,7 +6,7 @@
from collections.abc import Callable
from datetime import datetime
from enum import StrEnum
from typing import Optional
from typing import Optional, cast

import httpx

Expand All @@ -20,6 +20,7 @@
BetaContentBlockParam,
BetaMessageParam,
BetaTextBlockParam,
BetaToolResultBlockParam,
)

from .message_handler import MessageBuilder, ResponseProcessor
Expand Down Expand Up @@ -99,6 +100,12 @@ async def sampling_loop(
)

while True:
if only_n_most_recent_images:
_maybe_filter_to_n_most_recent_images(
messages,
only_n_most_recent_images,
)

# Configure client and betas
betas = []
enable_prompt_caching = False
Expand Down Expand Up @@ -176,34 +183,77 @@ async def sampling_loop(
def _maybe_filter_to_n_most_recent_images(
messages: list[BetaMessageParam],
images_to_keep: int,
min_removal_threshold: int = 10,
min_removal_threshold: int = 1,
):
"""
Filter messages to keep only the N most recent images.
With the assumption that images are screenshots that are of diminishing value as
the conversation progresses, remove all but the final `images_to_keep` tool_result
images and user images in place.
"""
if images_to_keep <= 0:
raise ValueError("images_to_keep must be > 0")
if images_to_keep is None or images_to_keep <= 0:
return

tool_result_blocks = cast(
list[BetaToolResultBlockParam],
[
item
for message in messages
for item in (
message["content"] if isinstance(message.get("content"), list) else []
)
if isinstance(item, dict) and item.get("type") == "tool_result"
],
)

total_images = sum(
1
for tool_result in tool_result_blocks
for content in (
tool_result.get("content")
if isinstance(tool_result.get("content"), list)
else []
)
if isinstance(content, dict) and content.get("type") == "image"
)

# Also count direct top-level image blocks in user messages if any
total_images += sum(
1
for message in messages
if message["role"] == "user"
for block in message.get("content", [])
if message.get("role") == "user" and isinstance(message.get("content"), list)
for block in message["content"]
if isinstance(block, dict) and block.get("type") == "image"
)

images_to_remove = total_images - images_to_keep
if images_to_remove < min_removal_threshold:
if min_removal_threshold > 1:
images_to_remove -= images_to_remove % min_removal_threshold

if images_to_remove <= 0:
return

images_removed = 0
for message in messages:
if message["role"] == "user" and isinstance(message.get("content"), list):
if message.get("role") == "user" and isinstance(message.get("content"), list):
new_content = []
for block in message["content"]:
if isinstance(block, dict) and block.get("type") == "image":
if images_removed < images_to_remove:
images_removed += 1
continue
if isinstance(block, dict):
if block.get("type") == "image":
if images_to_remove > 0:
images_to_remove -= 1
continue
elif block.get("type") == "tool_result" and isinstance(
block.get("content"), list
):
new_tool_content = []
for sub_block in block["content"]:
if (
isinstance(sub_block, dict)
and sub_block.get("type") == "image"
):
if images_to_remove > 0:
images_to_remove -= 1
continue
new_tool_content.append(sub_block)
block["content"] = new_tool_content
new_content.append(block)
message["content"] = new_content
21 changes: 15 additions & 6 deletions browser-use-demo/browser_use_demo/streamlit.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@ def setup_state():
"messages": [],
"system_prompt": "",
"hide_screenshots": False,
"only_n_most_recent_images": 3,
"rendered_message_count": 0, # Track rendered messages to avoid re-rendering
"last_error": None, # Store last error message to display persistently
# API Configuration
Expand Down Expand Up @@ -497,7 +498,7 @@ def api_response_callback(request, response, error):
api_key=st.session_state.api_key,
max_tokens=st.session_state.max_tokens,
browser_tool=st.session_state.browser_tool, # Pass persistent browser instance
only_n_most_recent_images=3, # Keep only 3 most recent screenshots for context
only_n_most_recent_images=st.session_state.only_n_most_recent_images,
)

# Update session state with the complete message history
Expand All @@ -523,21 +524,20 @@ def api_response_callback(request, response, error):
st.session_state.last_error = {"message": error_msg, "traceback": error_traceback}
with st.session_state.active_response_container:
st.error(error_msg)
st.code(error_traceback)
st.session_state.chat_disabled = False
st.rerun()


def main():
"""Main application entry point."""
# Set page configuration
st.set_page_config(
page_title="Claude Browser Use Demo",
page_title="Browser Use Demo",
page_icon="🌐",
layout="wide"
layout="wide",
initial_sidebar_state="expanded",
)

st.markdown(STREAMLIT_STYLE, unsafe_allow_html=True)

setup_state()


Expand Down Expand Up @@ -585,6 +585,15 @@ def main():
help="Add custom instructions for the browser agent",
)

# Only send N most recent images
st.number_input(
"Only send N most recent images",
min_value=0,
value=st.session_state.only_n_most_recent_images,
key="only_n_most_recent_images",
help="To decrease the total tokens sent, remove older screenshots from the conversation",
)

# Hide screenshots
st.checkbox(
"Hide Screenshots",
Expand Down
Loading