migrate memory banks to Resource and new registration

This commit is contained in:
Dinesh Yeduguru 2024-11-08 15:45:26 -08:00
parent b4416b72fd
commit c82f13bf9e
16 changed files with 178 additions and 104 deletions

View file

@ -5,7 +5,6 @@
# the root directory of this source tree.
import asyncio
import json
from typing import Any, Dict, List, Optional
@ -26,13 +25,13 @@ def deserialize_memory_bank_def(
raise ValueError("Memory bank type not specified")
type = j["type"]
if type == MemoryBankType.vector.value:
return VectorMemoryBankDef(**j)
return VectorMemoryBank(**j)
elif type == MemoryBankType.keyvalue.value:
return KeyValueMemoryBankDef(**j)
return KeyValueMemoryBank(**j)
elif type == MemoryBankType.keyword.value:
return KeywordMemoryBankDef(**j)
return KeywordMemoryBank(**j)
elif type == MemoryBankType.graph.value:
return GraphMemoryBankDef(**j)
return GraphMemoryBank(**j)
else:
raise ValueError(f"Unknown memory bank type: {type}")
@ -47,7 +46,7 @@ class MemoryBanksClient(MemoryBanks):
async def shutdown(self) -> None:
pass
async def list_memory_banks(self) -> List[MemoryBankDefWithProvider]:
async def list_memory_banks(self) -> List[MemoryBank]:
async with httpx.AsyncClient() as client:
response = await client.get(
f"{self.base_url}/memory_banks/list",
@ -57,13 +56,20 @@ class MemoryBanksClient(MemoryBanks):
return [deserialize_memory_bank_def(x) for x in response.json()]
async def register_memory_bank(
self, memory_bank: MemoryBankDefWithProvider
self,
memory_bank_id: str,
memory_bank_type: MemoryBankType,
provider_resource_id: Optional[str] = None,
provider_id: Optional[str] = None,
) -> None:
async with httpx.AsyncClient() as client:
response = await client.post(
f"{self.base_url}/memory_banks/register",
json={
"memory_bank": json.loads(memory_bank.json()),
"memory_bank_id": memory_bank_id,
"memory_bank_type": memory_bank_type.value,
"provider_resource_id": provider_resource_id,
"provider_id": provider_id,
},
headers={"Content-Type": "application/json"},
)
@ -71,13 +77,13 @@ class MemoryBanksClient(MemoryBanks):
async def get_memory_bank(
self,
identifier: str,
) -> Optional[MemoryBankDefWithProvider]:
memory_bank_id: str,
) -> Optional[MemoryBank]:
async with httpx.AsyncClient() as client:
response = await client.get(
f"{self.base_url}/memory_banks/get",
params={
"identifier": identifier,
"memory_bank_id": memory_bank_id,
},
headers={"Content-Type": "application/json"},
)
@ -94,7 +100,7 @@ async def run_main(host: str, port: int, stream: bool):
# register memory bank for the first time
response = await client.register_memory_bank(
VectorMemoryBankDef(
VectorMemoryBank(
identifier="test_bank2",
embedding_model="all-MiniLM-L6-v2",
chunk_size_in_tokens=512,

View file

@ -8,8 +8,10 @@ from enum import Enum
from typing import List, Literal, Optional, Protocol, runtime_checkable, Union
from llama_models.schema_utils import json_schema_type, webmethod
from pydantic import BaseModel, Field
from typing_extensions import Annotated
from pydantic import BaseModel
from llama_stack.apis.resource import Resource, ResourceType
@json_schema_type
@ -20,59 +22,121 @@ class MemoryBankType(Enum):
graph = "graph"
class CommonDef(BaseModel):
identifier: str
# Hack: move this out later
provider_id: str = ""
@json_schema_type
class MemoryBank(Resource):
type: Literal[ResourceType.memory_bank.value] = ResourceType.memory_bank.value
memory_bank_type: MemoryBankType
@json_schema_type
class VectorMemoryBankDef(CommonDef):
type: Literal[MemoryBankType.vector.value] = MemoryBankType.vector.value
class VectorMemoryBank(MemoryBank):
memory_bank_type: Literal[MemoryBankType.vector.value] = MemoryBankType.vector.value
embedding_model: str
chunk_size_in_tokens: int
overlap_size_in_tokens: Optional[int] = None
@json_schema_type
class KeyValueMemoryBankDef(CommonDef):
type: Literal[MemoryBankType.keyvalue.value] = MemoryBankType.keyvalue.value
class KeyValueMemoryBank(MemoryBank):
memory_bank_type: Literal[MemoryBankType.keyvalue.value] = (
MemoryBankType.keyvalue.value
)
@json_schema_type
class KeywordMemoryBankDef(CommonDef):
type: Literal[MemoryBankType.keyword.value] = MemoryBankType.keyword.value
class KeywordMemoryBank(MemoryBank):
memory_bank_type: Literal[MemoryBankType.keyword.value] = (
MemoryBankType.keyword.value
)
@json_schema_type
class GraphMemoryBankDef(CommonDef):
type: Literal[MemoryBankType.graph.value] = MemoryBankType.graph.value
class GraphMemoryBank(MemoryBank):
memory_bank_type: Literal[MemoryBankType.graph.value] = MemoryBankType.graph.value
MemoryBankDef = Annotated[
Union[
VectorMemoryBankDef,
KeyValueMemoryBankDef,
KeywordMemoryBankDef,
GraphMemoryBankDef,
],
Field(discriminator="type"),
@json_schema_type
class BaseRegistration(BaseModel):
memory_bank_id: str
provider_resource_id: Optional[str] = None
provider_id: Optional[str] = None
@json_schema_type
class VectorRegistration(BaseRegistration):
embedding_model: str
chunk_size_in_tokens: int
overlap_size_in_tokens: Optional[int] = None
@json_schema_type
class KeyValueRegistration(BaseRegistration):
pass
@json_schema_type
class KeywordRegistration(BaseRegistration):
pass
@json_schema_type
class GraphRegistration(BaseRegistration):
pass
RegistrationRequest = Union[
VectorRegistration,
KeyValueRegistration,
KeywordRegistration,
GraphRegistration,
]
MemoryBankDefWithProvider = MemoryBankDef
def registration_request_to_memory_bank(request: RegistrationRequest) -> MemoryBank:
"""Convert registration request to memory bank object"""
if isinstance(request, VectorRegistration):
return VectorMemoryBank(
identifier=request.memory_bank_id,
provider_resource_id=request.provider_resource_id,
provider_id=request.provider_id,
embedding_model=request.embedding_model,
chunk_size_in_tokens=request.chunk_size_in_tokens,
overlap_size_in_tokens=request.overlap_size_in_tokens,
)
elif isinstance(request, KeyValueRegistration):
return KeyValueMemoryBank(
identifier=request.memory_bank_id,
provider_resource_id=request.provider_resource_id,
provider_id=request.provider_id,
memory_bank_type=MemoryBankType.keyvalue,
)
elif isinstance(request, KeywordRegistration):
return KeywordMemoryBank(
identifier=request.memory_bank_id,
provider_resource_id=request.provider_resource_id,
provider_id=request.provider_id,
memory_bank_type=MemoryBankType.keyword,
)
elif isinstance(request, GraphRegistration):
return GraphMemoryBank(
identifier=request.memory_bank_id,
provider_resource_id=request.provider_resource_id,
provider_id=request.provider_id,
memory_bank_type=MemoryBankType.graph,
)
else:
raise ValueError(f"Unknown registration type: {type(request)}")
@runtime_checkable
class MemoryBanks(Protocol):
@webmethod(route="/memory_banks/list", method="GET")
async def list_memory_banks(self) -> List[MemoryBankDefWithProvider]: ...
async def list_memory_banks(self) -> List[MemoryBank]: ...
@webmethod(route="/memory_banks/get", method="GET")
async def get_memory_bank(
self, identifier: str
) -> Optional[MemoryBankDefWithProvider]: ...
async def get_memory_bank(self, memory_bank_id: str) -> Optional[MemoryBank]: ...
@webmethod(route="/memory_banks/register", method="POST")
async def register_memory_bank(
self, memory_bank: MemoryBankDefWithProvider
) -> None: ...
self, request: RegistrationRequest
) -> MemoryBank: ...