Compare commits

...

5 Commits

Author SHA1 Message Date
Nuno Campos 020b5941ab Lint 2023-11-05 19:56:11 +00:00
Nuno Campos 90b76f6c93 Fix unpack_input 2023-11-05 18:57:51 +00:00
Nuno Campos b70a63fb2f Update return values of storage endpoints 2023-11-05 11:25:10 +00:00
Nuno Campos 5dd75e2713 Less duplication 2023-11-05 11:25:10 +00:00
Nuno Campos ab589dac01 Implement storage for configs 2023-11-05 11:25:10 +00:00
5 changed files with 214 additions and 18 deletions
+5 -2
View File
@@ -6,6 +6,7 @@ This example shows how to use two options for configuration of runnables:
1) Configurable Fields: Use this to specify values for a given initialization parameter
2) Configurable Alternatives: Use this to specify complete alternative runnables
"""
import aiosqlite
from fastapi import FastAPI
from fastapi.middleware.cors import CORSMiddleware
from langchain.chat_models import ChatOpenAI
@@ -57,9 +58,11 @@ app.add_middleware(
# Add routes requires you to specify which config keys are accepted
# specifically, you must accept `configurable` as a config key.
add_routes(app, chain, config_keys=["configurable"])
add_routes(
app, chain, config_keys=["configurable"], storage=aiosqlite.connect(":memory:")
)
if __name__ == "__main__":
import uvicorn
uvicorn.run(app, host="localhost", port=8000)
uvicorn.run(app, host="localhost", port=8100)
+163 -15
View File
@@ -24,6 +24,7 @@ from typing import (
Union,
)
import aiosqlite
from fastapi import HTTPException, Request
from fastapi.exceptions import RequestValidationError
from langchain.callbacks.tracers.log_stream import RunLogPatch
@@ -49,6 +50,7 @@ try:
except ImportError:
from pydantic import BaseModel, Field, ValidationError, create_model
import langserve.storage
from langserve.playground import serve_playground
from langserve.serialization import WellKnownLCSerializer
from langserve.validation import (
@@ -87,13 +89,15 @@ def _config_from_hash(config_hash: str) -> Dict[str, Any]:
def _unpack_request_config(
*configs: Union[BaseModel, Mapping, str],
*configs: Union[BaseModel, Mapping, str, None],
keys: Sequence[str],
model: Type[BaseModel],
) -> Dict[str, Any]:
"""Merge configs, and project the given keys from the merged dict."""
config_dicts = []
for config in configs:
if config is None:
continue
if isinstance(config, str):
config_dicts.append(model(**_config_from_hash(config)).dict())
elif isinstance(config, BaseModel):
@@ -120,7 +124,10 @@ def _unpack_input(validated_model: BaseModel) -> Any:
# it was created by the server as part of validation and isn't expected
# to be accepted by the runnables as input as a pydantic model,
# instead we need to convert it into a corresponding python dict.
return model.dict()
return {
fieldname: _unpack_input(getattr(model, fieldname))
for fieldname in model.__fields__.keys()
}
return model
@@ -330,6 +337,7 @@ def add_routes(
config_keys: Sequence[str] = (),
include_callback_events: bool = False,
enable_feedback_endpoint: bool = False,
storage: Optional[aiosqlite.Connection] = None,
) -> None:
"""Register the routes on the given FastAPI app or APIRouter.
@@ -442,6 +450,27 @@ def add_routes(
if with_types:
runnable = runnable.with_types(**with_types)
if storage:
@app.get(f"{namespace}/s/")
async def list_stored_configs() -> Any:
"""Return a list of stored configs."""
return await langserve.storage.list_configs(storage)
@app.put(f"{namespace}/s/")
async def store_config(config: dict) -> Any:
"""Store a config."""
return await langserve.storage.set_config(
storage, config["key"], config["config"]
)
@app.on_event("shutdown")
async def shutdown_event():
try:
await storage.close()
except Exception:
pass
input_type_ = _resolve_model(runnable.get_input_schema(), "Input", model_namespace)
output_type_ = _resolve_model(
@@ -471,7 +500,7 @@ def add_routes(
BatchResponse = create_batch_response_model(model_namespace, output_type_)
async def _get_config_and_input(
request: Request, config_hash: str
request: Request, config_hash: str, config_id: str
) -> Tuple[RunnableConfig, Any]:
"""Extract the config and input from the request, validating the request."""
try:
@@ -484,6 +513,7 @@ def add_routes(
# Merge the config from the path with the config from the body.
config = _unpack_request_config(
config_hash,
await langserve.storage.get_config(storage, config_id),
body.config,
keys=config_keys,
model=ConfigPayload,
@@ -508,11 +538,12 @@ def add_routes(
async def invoke(
request: Request,
config_hash: str = "",
config_id: str = "",
) -> InvokeResponse:
"""Invoke the runnable with the given input and config."""
# We do not use the InvokeRequest model here since configurable runnables
# have dynamic schema -- so the validation below is a bit more involved.
config, input_ = await _get_config_and_input(request, config_hash)
config, input_ = await _get_config_and_input(request, config_hash, config_id)
event_aggregator = AsyncEventAggregatorCallback()
_add_tracing_info_to_metadata(config, request)
@@ -537,6 +568,13 @@ def add_routes(
),
)
if storage:
app.post(
namespace + "/s/{config_id}/invoke",
response_model=InvokeResponse,
include_in_schema=False,
)(invoke)
@app.post(
namespace + "/c/{config_hash}/batch",
response_model=BatchResponse,
@@ -548,6 +586,7 @@ def add_routes(
async def batch(
request: Request,
config_hash: str = "",
config_id: str = "",
) -> BatchResponse:
"""Invoke the runnable with the given inputs and config."""
try:
@@ -555,6 +594,8 @@ def add_routes(
except json.JSONDecodeError:
raise RequestValidationError(errors=["Invalid JSON body"])
config_from_id = await langserve.storage.get_config(storage, config_id)
with _with_validation_error_translation():
body = BatchRequestShallowValidator.validate(body)
config = body.config
@@ -570,13 +611,18 @@ def add_routes(
configs = [
_unpack_request_config(
config_hash, config, keys=config_keys, model=ConfigPayload
config_hash,
config_from_id,
config,
keys=config_keys,
model=ConfigPayload,
)
for config in config
]
elif isinstance(config, dict):
configs = _unpack_request_config(
config_hash,
config_from_id,
config,
keys=config_keys,
model=ConfigPayload,
@@ -631,11 +677,19 @@ def add_routes(
),
)
if storage:
app.post(
namespace + "/s/{config_id}/batch",
response_model=BatchResponse,
include_in_schema=False,
)(batch)
@app.post(namespace + "/c/{config_hash}/stream", include_in_schema=False)
@app.post(f"{namespace}/stream", include_in_schema=False)
async def stream(
request: Request,
config_hash: str = "",
config_id: str = "",
) -> EventSourceResponse:
"""Invoke the runnable stream the output.
@@ -645,7 +699,9 @@ def add_routes(
err_event = {}
validation_exception: Optional[BaseException] = None
try:
config, input_ = await _get_config_and_input(request, config_hash)
config, input_ = await _get_config_and_input(
request, config_hash, config_id
)
except BaseException as e:
validation_exception = e
if isinstance(e, RequestValidationError):
@@ -700,6 +756,12 @@ def add_routes(
return EventSourceResponse(_stream())
if storage:
app.post(
namespace + "/s/{config_id}/stream",
include_in_schema=False,
)(stream)
@app.post(
namespace + "/c/{config_hash}/stream_log",
include_in_schema=False,
@@ -711,6 +773,7 @@ def add_routes(
async def stream_log(
request: Request,
config_hash: str = "",
config_id: str = "",
) -> EventSourceResponse:
"""Invoke the runnable stream_log the output.
@@ -720,7 +783,9 @@ def add_routes(
err_event = {}
validation_exception: Optional[BaseException] = None
try:
config, input_ = await _get_config_and_input(request, config_hash)
config, input_ = await _get_config_and_input(
request, config_hash, config_id
)
except BaseException as e:
validation_exception = e
if isinstance(e, RequestValidationError):
@@ -811,6 +876,12 @@ def add_routes(
return EventSourceResponse(_stream_log())
if storage:
app.post(
namespace + "/s/{config_id}/stream_log",
include_in_schema=False,
)(stream_log)
@app.get(
namespace + "/c/{config_hash}/input_schema",
tags=route_tags_with_config,
@@ -819,15 +890,25 @@ def add_routes(
@app.get(
f"{namespace}/input_schema", tags=route_tags, name=_route_name("input_schema")
)
async def input_schema(config_hash: str = "") -> Any:
async def input_schema(config_hash: str = "", config_id: str = "") -> Any:
"""Return the input schema of the runnable."""
with _with_validation_error_translation():
config = _unpack_request_config(
config_hash, keys=config_keys, model=ConfigPayload
config_hash,
await langserve.storage.get_config(storage, config_id),
keys=config_keys,
model=ConfigPayload,
)
return runnable.get_input_schema(config).schema()
if storage:
app.get(
namespace + "/s/{config_id}/input_schema",
tags=route_tags_with_config,
name=_route_name_with_config("input_schema"),
)(input_schema)
@app.get(
namespace + "/c/{config_hash}/output_schema",
tags=route_tags_with_config,
@@ -836,14 +917,24 @@ def add_routes(
@app.get(
f"{namespace}/output_schema", tags=route_tags, name=_route_name("output_schema")
)
async def output_schema(config_hash: str = "") -> Any:
async def output_schema(config_hash: str = "", config_id: str = "") -> Any:
"""Return the output schema of the runnable."""
with _with_validation_error_translation():
config = _unpack_request_config(
config_hash, keys=config_keys, model=ConfigPayload
config_hash,
await langserve.storage.get_config(storage, config_id),
keys=config_keys,
model=ConfigPayload,
)
return runnable.get_output_schema(config).schema()
if storage:
app.get(
namespace + "/s/{config_id}/output_schema",
tags=route_tags_with_config,
name=_route_name_with_config("output_schema"),
)(output_schema)
@app.get(
namespace + "/c/{config_hash}/config_schema",
tags=route_tags_with_config,
@@ -852,24 +943,39 @@ def add_routes(
@app.get(
f"{namespace}/config_schema", tags=route_tags, name=_route_name("config_schema")
)
async def config_schema(config_hash: str = "") -> Any:
async def config_schema(config_hash: str = "", config_id: str = "") -> Any:
"""Return the config schema of the runnable."""
with _with_validation_error_translation():
config = _unpack_request_config(
config_hash, keys=config_keys, model=ConfigPayload
config_hash,
await langserve.storage.get_config(storage, config_id),
keys=config_keys,
model=ConfigPayload,
)
return runnable.with_config(config).config_schema(include=config_keys).schema()
if storage:
app.get(
namespace + "/s/{config_id}/config_schema",
tags=route_tags_with_config,
name=_route_name_with_config("config_schema"),
)(config_schema)
@app.get(
namespace + "/c/{config_hash}/playground/{file_path:path}",
include_in_schema=False,
)
@app.get(namespace + "/playground/{file_path:path}", include_in_schema=False)
async def playground(file_path: str, config_hash: str = "") -> Any:
async def playground(
file_path: str, config_hash: str = "", config_id: str = ""
) -> Any:
"""Return the playground of the runnable."""
with _with_validation_error_translation():
config = _unpack_request_config(
config_hash, keys=config_keys, model=ConfigPayload
config_hash,
await langserve.storage.get_config(storage, config_id),
keys=config_keys,
model=ConfigPayload,
)
return await serve_playground(
runnable.with_config(config),
@@ -879,6 +985,12 @@ def add_routes(
file_path,
)
if storage:
app.get(
namespace + "/s/{config_id}/playground/{file_path:path}",
include_in_schema=False,
)(playground)
@app.post(namespace + "/feedback")
async def feedback(feedback_create_req: FeedbackCreateRequest) -> Feedback:
"""
@@ -935,10 +1047,19 @@ def add_routes(
async def _invoke_docs(
invoke_request: Annotated[InvokeRequest, InvokeRequest],
config_hash: str = "",
config_id: str = "",
) -> InvokeResponse:
"""Invoke the runnable with the given input and config."""
raise AssertionError("This endpoint should not be reachable.")
if storage:
app.post(
namespace + "/s/{config_id}/invoke",
response_model=InvokeResponse,
tags=route_tags_with_config,
name=_route_name_with_config("invoke"),
)(_invoke_docs)
@app.post(
namespace + "/c/{config_hash}/batch",
response_model=BatchResponse,
@@ -954,10 +1075,19 @@ def add_routes(
async def _batch_docs(
batch_request: Annotated[BatchRequest, BatchRequest],
config_hash: str = "",
config_id: str = "",
) -> BatchResponse:
"""Batch invoke the runnable with the given inputs and config."""
raise AssertionError("This endpoint should not be reachable.")
if storage:
app.post(
namespace + "/s/{config_id}/batch",
response_model=BatchResponse,
tags=route_tags_with_config,
name=_route_name_with_config("batch"),
)(_batch_docs)
@app.post(
namespace + "/c/{config_hash}/stream",
include_in_schema=True,
@@ -973,6 +1103,7 @@ def add_routes(
async def _stream_docs(
stream_request: Annotated[StreamRequest, StreamRequest],
config_hash: str = "",
config_id: str = "",
) -> EventSourceResponse:
"""Invoke the runnable stream the output.
@@ -1017,6 +1148,14 @@ def add_routes(
"""
raise AssertionError("This endpoint should not be reachable.")
if storage:
app.post(
namespace + "/s/{config_id}/stream",
include_in_schema=True,
tags=route_tags_with_config,
name=_route_name_with_config("stream"),
)(_stream_docs)
@app.post(
namespace + "/c/{config_hash}/stream_log",
include_in_schema=True,
@@ -1032,6 +1171,7 @@ def add_routes(
async def _stream_log_docs(
stream_log_request: Annotated[StreamLogRequest, StreamLogRequest],
config_hash: str = "",
config_id: str = "",
) -> EventSourceResponse:
"""Invoke the runnable stream_log the output.
@@ -1076,3 +1216,11 @@ def add_routes(
}
"""
raise AssertionError("This endpoint should not be reachable.")
if storage:
app.post(
namespace + "/s/{config_id}/stream_log",
include_in_schema=True,
tags=route_tags_with_config,
name=_route_name_with_config("stream_log"),
)(_stream_log_docs)
+43
View File
@@ -0,0 +1,43 @@
import json
import weakref
import aiosqlite
SEEN = weakref.WeakSet()
async def init(db: aiosqlite.Connection):
if db in SEEN:
return
SEEN.add(db)
await db
db.row_factory = aiosqlite.Row
await db.execute(
"CREATE TABLE IF NOT EXISTS configs (key TEXT PRIMARY KEY, config TEXT)"
)
await db.commit()
async def list_configs(db: aiosqlite.Connection):
await init(db)
cursor = await db.execute("SELECT * FROM configs")
return await cursor.fetchall()
async def get_config(db: aiosqlite.Connection, key: str):
if not key or not db:
return None
await init(db)
cursor = await db.execute("SELECT config FROM configs WHERE key=?", (key,))
row = await cursor.fetchone()
return json.loads(row[0]) if row else None
async def set_config(db: aiosqlite.Connection, key: str, config: dict):
await init(db)
await db.execute(
"INSERT OR REPLACE INTO configs VALUES (?, ?)", (key, json.dumps(config))
)
await db.commit()
return {"key": key, "config": config}
Generated
+1 -1
View File
@@ -3822,4 +3822,4 @@ server = ["fastapi", "sse-starlette"]
[metadata]
lock-version = "2.0"
python-versions = "^3.8.1"
content-hash = "5175284ed862c0a87cb8cbe166795841069ad416c8e1549fcd9974c1a4d4fe5f"
content-hash = "dcc245797a409da9fec1f58961b1a39edf8cbc6129284f89f29bb09b579f2127"
+2
View File
@@ -17,12 +17,14 @@ sse-starlette = {version = "^1.3.0", optional = true}
httpx-sse = {version = ">=0.3.1", optional = true}
pydantic = "^1"
langchain = ">=0.0.322"
aiosqlite = {version = "^0.19.0", optional = true}
[tool.poetry.group.dev.dependencies]
jupyterlab = "^3.6.1"
fastapi = ">=0.90.1"
sse-starlette = "^1.3.0"
httpx-sse = ">=0.3.1"
aiosqlite = "^0.19.0"
[tool.poetry.group.typing.dependencies]