refactor objectmodel

This commit is contained in:
LUIS NOVO 2024-11-19 19:03:32 -03:00
parent f140a5e228
commit c297dcb809
8 changed files with 186 additions and 68 deletions

View file

@ -40,6 +40,11 @@ def repo_create(table: str, data: Dict[str, Any]):
return repo_query(query) return repo_query(query)
def repo_upsert(table: str, data: Dict[str, Any]):
query = f"UPSERT {table} CONTENT {data};"
return repo_query(query)
def repo_update(id: str, data: Dict[str, Any]): def repo_update(id: str, data: Dict[str, Any]):
query = "UPDATE $id CONTENT $data;" query = "UPDATE $id CONTENT $data;"
vars = {"id": id, "data": data} vars = {"id": id, "data": data}

View file

@ -1,8 +1,22 @@
from datetime import datetime from datetime import datetime
from typing import Any, ClassVar, Dict, List, Optional, Type, TypeVar, cast from typing import (
Any,
ClassVar,
Dict,
List,
Optional,
Type,
TypeVar,
cast,
)
from loguru import logger from loguru import logger
from pydantic import BaseModel, ValidationError, field_validator from pydantic import (
BaseModel,
ValidationError,
field_validator,
model_validator,
)
from open_notebook.database.repository import ( from open_notebook.database.repository import (
repo_create, repo_create,
@ -10,6 +24,7 @@ from open_notebook.database.repository import (
repo_query, repo_query,
repo_relate, repo_relate,
repo_update, repo_update,
repo_upsert,
) )
from open_notebook.exceptions import ( from open_notebook.exceptions import (
DatabaseOperationError, DatabaseOperationError,
@ -204,24 +219,92 @@ class ObjectModel(BaseModel):
class RecordModel(BaseModel): class RecordModel(BaseModel):
record_id: ClassVar[str] record_id: ClassVar[str]
auto_save: ClassVar[bool] = (
False # Default to False, can be overridden in subclasses
)
_instances: ClassVar[Dict[str, "RecordModel"]] = {} # Store instances by record_id
class Config:
validate_assignment = True
arbitrary_types_allowed = True
extra = "allow"
from_attributes = True
defer_build = True
def __new__(cls, **kwargs):
# If an instance already exists for this record_id, return it
if cls.record_id in cls._instances:
instance = cls._instances[cls.record_id]
# Update instance with any new kwargs if provided
if kwargs:
for key, value in kwargs.items():
setattr(instance, key, value)
return instance
# If no instance exists, create a new one
instance = super().__new__(cls)
cls._instances[cls.record_id] = instance
return instance
def __init__(self, **kwargs): def __init__(self, **kwargs):
super().__init__(**kwargs) # Only initialize if this is a new instance
self.load() if not hasattr(self, "_initialized"):
object.__setattr__(self, "__dict__", {})
# Load data from DB first
result = repo_query(f"SELECT * FROM {self.record_id};")
if result:
db_data = result[0]
else:
# Initialize empty object with None for Optional fields
db_data = {
field_name: None
for field_name, field_info in self.model_fields.items()
if not str(field_info.annotation).startswith("typing.ClassVar")
}
# Initialize with DB data and any overrides
super().__init__(**{**db_data, **kwargs})
object.__setattr__(self, "_initialized", True)
@classmethod
def get_instance(cls) -> "RecordModel":
"""Get or create the singleton instance"""
return cls()
@model_validator(mode="after")
def auto_save_validator(self):
if self.__class__.auto_save:
self.update()
return self
def update(self):
# Get all non-ClassVar fields and their values
data = {
field_name: getattr(self, field_name)
for field_name, field_info in self.model_fields.items()
if not str(field_info.annotation).startswith("typing.ClassVar")
}
repo_upsert(self.record_id, data)
def load(self):
result = repo_query(f"SELECT * FROM {self.record_id};") result = repo_query(f"SELECT * FROM {self.record_id};")
if result: if result:
result = result[0] for key, value in result[0].items():
else: if hasattr(self, key):
repo_create(self.record_id, {}) object.__setattr__(
result = {} self, key, value
for key, value in result.items(): ) # Use object.__setattr__ to avoid triggering validation again
if hasattr(self, key):
setattr(self, key, value)
return self return self
def update(self, data): @classmethod
repo_update(self.record_id, data) def clear_instance(cls):
return self.load() """Clear the singleton instance (useful for testing)"""
if cls.record_id in cls._instances:
del cls._instances[cls.record_id]
def patch(self, model_dict: dict):
"""Update model attributes from dictionary and save"""
for key, value in model_dict.items():
setattr(self, key, value)
self.update()

View file

@ -28,15 +28,14 @@ class Model(ObjectModel):
class DefaultModels(RecordModel): class DefaultModels(RecordModel):
record_id: ClassVar[str] = "open_notebook:default_models" record_id: ClassVar[str] = "open_notebook:default_models"
default_chat_model: Optional[str]
default_chat_model: Optional[str] = None default_transformation_model: Optional[str]
default_transformation_model: Optional[str] = None large_context_model: Optional[str]
large_context_model: Optional[str] = None default_text_to_speech_model: Optional[str]
default_text_to_speech_model: Optional[str] = None default_speech_to_text_model: Optional[str]
default_speech_to_text_model: Optional[str] = None # default_vision_model: Optional[str]
# default_vision_model: Optional[str] = None default_embedding_model: Optional[str]
default_embedding_model: Optional[str] = None default_tools_model: Optional[str]
default_tools_model: Optional[str] = None
class ModelManager: class ModelManager:
@ -54,7 +53,10 @@ class ModelManager:
self._default_models = None self._default_models = None
self.refresh_defaults() self.refresh_defaults()
def get_model(self, model_id: str, **kwargs) -> ModelType: def get_model(self, model_id: str, **kwargs) -> Optional[ModelType]:
if not model_id:
return None
cache_key = f"{model_id}:{str(kwargs)}" cache_key = f"{model_id}:{str(kwargs)}"
if cache_key in self._model_cache: if cache_key in self._model_cache:
@ -68,9 +70,6 @@ class ModelManager:
) )
return cached_model return cached_model
if not model_id:
return None
model: Model = Model.get(model_id) model: Model = Model.get(model_id)
if not model: if not model:
@ -111,7 +110,10 @@ class ModelManager:
@property @property
def speech_to_text(self, **kwargs) -> Optional[SpeechToTextModel]: def speech_to_text(self, **kwargs) -> Optional[SpeechToTextModel]:
"""Get the default speech-to-text model""" """Get the default speech-to-text model"""
model = self.get_default_model("speech_to_text", **kwargs) model_id = self.defaults.default_speech_to_text_model
if not model_id:
return None
model = self.get_model(model_id, **kwargs)
assert model is None or isinstance( assert model is None or isinstance(
model, SpeechToTextModel model, SpeechToTextModel
), f"Expected SpeechToTextModel but got {type(model)}" ), f"Expected SpeechToTextModel but got {type(model)}"
@ -120,7 +122,10 @@ class ModelManager:
@property @property
def text_to_speech(self, **kwargs) -> Optional[TextToSpeechModel]: def text_to_speech(self, **kwargs) -> Optional[TextToSpeechModel]:
"""Get the default text-to-speech model""" """Get the default text-to-speech model"""
model = self.get_default_model("text_to_speech", **kwargs) model_id = self.defaults.default_text_to_speech_model
if not model_id:
return None
model = self.get_model(model_id, **kwargs)
assert model is None or isinstance( assert model is None or isinstance(
model, TextToSpeechModel model, TextToSpeechModel
), f"Expected TextToSpeechModel but got {type(model)}" ), f"Expected TextToSpeechModel but got {type(model)}"
@ -129,13 +134,16 @@ class ModelManager:
@property @property
def embedding_model(self, **kwargs) -> Optional[EmbeddingModel]: def embedding_model(self, **kwargs) -> Optional[EmbeddingModel]:
"""Get the default embedding model""" """Get the default embedding model"""
model = self.get_default_model("embedding", **kwargs) model_id = self.defaults.default_embedding_model
if not model_id:
return None
model = self.get_model(model_id, **kwargs)
assert model is None or isinstance( assert model is None or isinstance(
model, EmbeddingModel model, EmbeddingModel
), f"Expected EmbeddingModel but got {type(model)}" ), f"Expected EmbeddingModel but got {type(model)}"
return model return model
def get_default_model(self, model_type: str, **kwargs) -> ModelType: def get_default_model(self, model_type: str, **kwargs) -> Optional[ModelType]:
""" """
Get the default model for a specific type. Get the default model for a specific type.
@ -165,6 +173,9 @@ class ModelManager:
elif model_type == "large_context": elif model_type == "large_context":
model_id = self.defaults.large_context_model model_id = self.defaults.large_context_model
if not model_id:
return None
return self.get_model(model_id, **kwargs) return self.get_model(model_id, **kwargs)
def clear_cache(self): def clear_cache(self):

View file

@ -1,4 +1,3 @@
from executing import Source
from langchain_core.messages import HumanMessage, SystemMessage from langchain_core.messages import HumanMessage, SystemMessage
from langchain_core.runnables import ( from langchain_core.runnables import (
RunnableConfig, RunnableConfig,
@ -6,6 +5,7 @@ from langchain_core.runnables import (
from langgraph.graph import END, START, StateGraph from langgraph.graph import END, START, StateGraph
from typing_extensions import TypedDict from typing_extensions import TypedDict
from open_notebook.domain.notebook import Source
from open_notebook.domain.transformation import DefaultPrompts, Transformation from open_notebook.domain.transformation import DefaultPrompts, Transformation
from open_notebook.graphs.utils import provision_langchain_model from open_notebook.graphs.utils import provision_langchain_model
from open_notebook.prompter import Prompter from open_notebook.prompter import Prompter
@ -26,7 +26,7 @@ def run_transformation(state: dict, config: RunnableConfig) -> dict:
if not content: if not content:
content = source.full_text content = source.full_text
transformation_prompt_text = transformation.prompt transformation_prompt_text = transformation.prompt
default_prompts: DefaultPrompts = DefaultPrompts().load() default_prompts: DefaultPrompts = DefaultPrompts()
if default_prompts.transformation_instructions: if default_prompts.transformation_instructions:
transformation_prompt_text = f"{default_prompts.transformation_instructions}\n\n{transformation_prompt_text}" transformation_prompt_text = f"{default_prompts.transformation_instructions}\n\n{transformation_prompt_text}"

View file

@ -54,7 +54,7 @@ with ask_tab:
"The LLM will answer your query based on the documents in your knowledge base. " "The LLM will answer your query based on the documents in your knowledge base. "
) )
question = st.text_input("Question", "") question = st.text_input("Question", "")
default_model = DefaultModels().load().default_chat_model default_model = DefaultModels().default_chat_model
strategy_model = model_selector( strategy_model = model_selector(
"Query Strategy Model", "Query Strategy Model",
"strategy_model", "strategy_model",

View file

@ -89,7 +89,7 @@ def generate_new_models(models, suggested_models):
return new_models return new_models
default_models = DefaultModels().model_dump() default_models = DefaultModels()
all_models = Model.get_all() all_models = Model.get_all()
with model_tab: with model_tab:
@ -176,82 +176,101 @@ with model_defaults_tab:
"In this section, you can select the default models to be used on the various content operations done by Open Notebook. Some of these can be overriden in the different modules." "In this section, you can select the default models to be used on the various content operations done by Open Notebook. Some of these can be overriden in the different modules."
) )
defs = {} defs = {}
defs["default_chat_model"] = model_selector( # Handle chat model selection
selected_model = model_selector(
"Default Chat Model", "Default Chat Model",
"default_chat_model", "default_chat_model",
selected_id=default_models.get("default_chat_model"), selected_id=default_models.default_chat_model,
help="This model will be used for chat.", help="This model will be used for chat.",
model_type="language", model_type="language",
) )
if selected_model:
default_models.default_chat_model = selected_model.id
st.divider() st.divider()
defs["default_transformation_model"] = model_selector( # Handle transformation model selection
selected_model = model_selector(
"Default Transformation Model", "Default Transformation Model",
"default_transformation_model", "default_transformation_model",
selected_id=default_models.get("default_transformation_model"), selected_id=default_models.default_transformation_model,
help="This model will be used for text transformations such as summaries, insights, etc.", help="This model will be used for text transformations such as summaries, insights, etc.",
model_type="language", model_type="language",
) )
if selected_model:
default_models.default_transformation_model = selected_model.id
st.caption("You can use a cheap model here like gpt-4o-mini, llama3, etc.") st.caption("You can use a cheap model here like gpt-4o-mini, llama3, etc.")
st.divider() st.divider()
defs["default_tools_model"] = model_selector(
# Handle tools model selection
selected_model = model_selector(
"Default Tools Model", "Default Tools Model",
"default_tools_model", "default_tools_model",
selected_id=default_models.get("default_tools_model"), selected_id=default_models.default_tools_model,
help="This model will be used for calling tools. Currently, it's best to use Open AI and Anthropic for this.", help="This model will be used for calling tools. Currently, it's best to use Open AI and Anthropic for this.",
model_type="language", model_type="language",
) )
if selected_model:
default_models.default_tools_model = selected_model.id
st.caption("Recommended to use a capable model here, like gpt-4o, claude, etc.") st.caption("Recommended to use a capable model here, like gpt-4o, claude, etc.")
st.divider() st.divider()
defs["large_context_model"] = model_selector(
# Handle large context model selection
selected_model = model_selector(
"Large Context Model", "Large Context Model",
"large_context_model", "large_context_model",
selected_id=default_models.get("large_context_model"), selected_id=default_models.large_context_model,
help="This model will be used for larger context generation -- recommended: Gemini", help="This model will be used for larger context generation -- recommended: Gemini",
model_type="language", model_type="language",
) )
if selected_model:
default_models.large_context_model = selected_model.id
st.caption("Recommended to use Gemini models for larger context processing") st.caption("Recommended to use Gemini models for larger context processing")
st.divider() st.divider()
defs["default_text_to_speech_model"] = model_selector(
# Handle text-to-speech model selection
selected_model = model_selector(
"Default Text to Speech Model", "Default Text to Speech Model",
"default_text_to_speech_model", "default_text_to_speech_model",
selected_id=default_models.get("default_text_to_speech_model"), selected_id=default_models.default_text_to_speech_model,
help="This is the default model for converting text to speech (podcasts, etc)", help="This is the default model for converting text to speech (podcasts, etc)",
model_type="text_to_speech", model_type="text_to_speech",
) )
st.caption("You can override this model on different podcasts") st.caption("You can override this model on different podcasts")
if selected_model:
default_models.default_text_to_speech_model = selected_model.id
st.divider() st.divider()
defs["default_speech_to_text_model"] = model_selector(
# Handle speech-to-text model selection
selected_model = model_selector(
"Default Speech to Text Model", "Default Speech to Text Model",
"default_speech_to_text_model", selected_id=default_models.default_speech_to_text_model,
selected_id=default_models.get("default_speech_to_text_model"),
help="This is the default model for converting speech to text (audio transcriptions, etc)", help="This is the default model for converting speech to text (audio transcriptions, etc)",
model_type="speech_to_text", model_type="speech_to_text",
key="default_speech_to_text_model",
) )
st.divider() if selected_model:
# defs["default_vision_model"] = ( default_models.default_speech_to_text_model = selected_model.id
# model_selector(
# "Default Speech to Text Model",
# "default_vision_model",
# selected_id=default_models.get("default_vision_model"),
# help="This is the default model for vision tasks",
# model_type="vision",
# ),
# )
defs["default_embedding_model"] = model_selector( st.divider()
# Handle embedding model selection
selected_model = model_selector(
"Default Speech to Text Model", "Default Speech to Text Model",
"default_embedding_model", "default_embedding_model",
selected_id=default_models.get("default_embedding_model"), selected_id=default_models.default_embedding_model,
help="This is the default model for embeddings (semantic search, etc)", help="This is the default model for embeddings (semantic search, etc)",
model_type="embedding", model_type="embedding",
) )
st.caption( if selected_model:
default_models.default_embedding_model = selected_model.id
st.warning(
"Caution: you cannot change the embedding model once there is embeddings or they will need to be regenerated" "Caution: you cannot change the embedding model once there is embeddings or they will need to be regenerated"
) )
for k, v in defs.items(): for k, v in defs.items():
if v: if v:
defs[k] = v.id defs[k] = v.id
DefaultModels().update(defs)
model_manager.refresh_defaults() if st.button("Save Defaults"):
default_models.patch(defs)
model_manager.refresh_defaults()
st.success("Saved")

View file

@ -25,7 +25,7 @@ with transformations_tab:
st.markdown( st.markdown(
"Transformations are prompts that will be used by the LLM to process a source and extract insights, summaries, etc. " "Transformations are prompts that will be used by the LLM to process a source and extract insights, summaries, etc. "
) )
default_prompts: DefaultPrompts = DefaultPrompts().load() default_prompts: DefaultPrompts = DefaultPrompts()
with st.expander("**⚙️ Default Transformation Prompt**"): with st.expander("**⚙️ Default Transformation Prompt**"):
default_prompts.transformation_instructions = st.text_area( default_prompts.transformation_instructions = st.text_area(
"Default Transformation Prompt", "Default Transformation Prompt",
@ -34,7 +34,7 @@ with transformations_tab:
) )
st.caption("This will be added to all your transformation prompts.") st.caption("This will be added to all your transformation prompts.")
if st.button("Save", key="save_default_prompt"): if st.button("Save", key="save_default_prompt"):
default_prompts.update(default_prompts.model_dump()) default_prompts.update()
st.toast("Default prompt saved successfully!") st.toast("Default prompt saved successfully!")
if st.button("Create new Transformation", icon="", key="new_transformation"): if st.button("Create new Transformation", icon="", key="new_transformation"):
new_transformation = Transformation( new_transformation = Transformation(

View file

@ -6,7 +6,7 @@ import streamlit as st
from loguru import logger from loguru import logger
from open_notebook.database.migrate import MigrationManager from open_notebook.database.migrate import MigrationManager
from open_notebook.domain.models import model_manager from open_notebook.domain.models import DefaultModels
from open_notebook.domain.notebook import ChatSession, Notebook from open_notebook.domain.notebook import ChatSession, Notebook
from open_notebook.graphs.chat import ThreadState, graph from open_notebook.graphs.chat import ThreadState, graph
from open_notebook.utils import ( from open_notebook.utils import (
@ -109,7 +109,7 @@ def check_migration():
def check_models(only_mandatory=True, stop_on_error=True): def check_models(only_mandatory=True, stop_on_error=True):
default_models = model_manager.defaults default_models = DefaultModels()
mandatory_models = [ mandatory_models = [
default_models.default_chat_model, default_models.default_chat_model,
default_models.default_transformation_model, default_models.default_transformation_model,