Streamlit 基于 asyncio(底层用 Tornado),但事件循环不暴露给用户代码。本文探讨如何在 Streamlit 中正确使用 asyncio 数据源。
Problem 1 2 3 4 5 async def fetch_data (): count = await Product.find(fetch_links=True ).count() if not count: await Product(name='test' ).save() return await Product.find(fetch_links=True ).to_list()
Solution 1: asyncio.run() 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 import asynciofrom beanie import Document, init_beaniefrom motor.motor_asyncio import AsyncIOMotorClientimport streamlit as stclass Product (Document ): name: str async def init_database (client ): database = client.get_database(name='asyncio_streamlit_db' ) await init_beanie(database=database, document_models=[Product])async def fetch_data (): count = await Product.find(fetch_links=True ).count() if not count: await Product(name='test' ).save() return await Product.find(fetch_links=True ).to_list()async def main (): client = AsyncIOMotorClient("mongodb://127.0.0.1:27017/" ) await init_database(client) products = await fetch_data() for product in products: st.write(product)if __name__ == '__main__' : asyncio.run(main())
问题 :每次请求都创建新连接,项目规模扩大后不现实。
Solution 2: Cache per session 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 import asynciofrom beanie import Document, init_beaniefrom motor.motor_asyncio import AsyncIOMotorClientimport streamlit as stclass Product (Document ): name: str async def init_database (client ): database = client.get_database(name='asyncio_streamlit_db' ) await init_beanie(database=database, document_models=[Product])async def fetch_data (): count = await Product.find(fetch_links=True ).count() if count < 10 : for i in range (10 ): await Product(name='test' ).save() return await Product.find(fetch_links=True ).limit(10 ).to_list()async def main (): if not st.session_state.get('client' ): st.session_state.client = AsyncIOMotorClient("mongodb://127.0.0.1:27017/" ) await init_database(st.session_state.client) products = await fetch_data() for product in products: st.write(product) st.button("Quick rerun" )if __name__ == '__main__' : asyncio.run(main())
问题 :按钮点击 rerun 时,asyncio.run() 创建新的事件循环,beanie 初始化状态丢失。
Solution 3: Cache event loop 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 import asynciofrom beanie import Document, init_beaniefrom motor.motor_asyncio import AsyncIOMotorClientimport streamlit as stclass Product (Document ): name: str async def init_database (client ): database = client.get_database(name='asyncio_streamlit_db' ) await init_beanie(database=database, document_models=[Product])async def fetch_data (): count = await Product.find(fetch_links=True ).count() if count < 10 : for i in range (10 ): await Product(name='test' ).save() return await Product.find(fetch_links=True ).limit(10 ).to_list()def get_event_loop (): return asyncio.new_event_loop()if not st.session_state.get('event_loop' ): st.session_state.event_loop = get_event_loop()if not st.session_state.get('client' ): st.session_state.client = AsyncIOMotorClient( "mongodb://127.0.0.1:27017/" , io_loop=st.session_state.event_loop )async def main (): await init_database(st.session_state.client) products = await fetch_data() for product in products: st.write(product) st.button("Quick rerun" )if __name__ == '__main__' : st.session_state.event_loop.run_until_complete(main())
问题 :快速点击按钮时 RuntimeError: This event loop is already running。
Solution 4: Worker thread (推荐) 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 import asynciofrom asyncio import run_coroutine_threadsafefrom threading import Threadfrom motor.motor_asyncio import AsyncIOMotorClientimport streamlit as st@st.cache_resource(show_spinner=False ) def create_loop (): loop = asyncio.new_event_loop() thread = Thread(target=loop.run_forever) thread.start() return loop, thread st.session_state.event_loop, worker_thread = create_loop()def run_async (coroutine ): return run_coroutine_threadsafe(coroutine, st.session_state.event_loop).result()@st.cache_resource(show_spinner=False ) def setup_database (): client = AsyncIOMotorClient( "mongodb://127.0.0.1:27017/" , io_loop=st.session_state.event_loop ) run_async(init_database(client=client)) return client st.session_state.db_client = setup_database()def main (): products = run_async(fetch_data()) for product in products: st.write(product) st.button("Quick rerun" )if __name__ == '__main__' : main()
优点 :独立的 worker 线程运行事件循环,run_coroutine_threadsafe 安全提交任务,st.cache_resource 缓存连接不重置。
缺点 :Streamlit 没有 shutdown 钩子,worker 线程无法优雅退出。
完整示例:https://github.com/thorin-schiffer/streamlit_asyncio