diff --git a/backend/alembic/versions/792d1af3dc44_create_table_user_slack_persona.py b/backend/alembic/versions/792d1af3dc44_create_table_user_slack_persona.py new file mode 100644 index 00000000000..ffbe61e4b53 --- /dev/null +++ b/backend/alembic/versions/792d1af3dc44_create_table_user_slack_persona.py @@ -0,0 +1,32 @@ +"""Create table user_slack_persona + +Revision ID: 792d1af3dc44 +Revises: 3a7802814195 +Create Date: 2025-01-24 04:26:02.844951 + +""" +from alembic import op +import sqlalchemy as sa + +# revision identifiers, used by Alembic. +revision = "792d1af3dc44" +down_revision = "3a7802814195" +branch_labels = None +depends_on = None + + +def upgrade() -> None: + op.create_table( + "user_slack_persona", + sa.Column("sender_id", sa.String(), nullable=False), + sa.Column("persona_id", sa.Integer(), nullable=True), + sa.ForeignKeyConstraint( + ["persona_id"], + ["persona.id"], + ), + sa.PrimaryKeyConstraint("sender_id"), + ) + + +def downgrade() -> None: + op.drop_table("user_slack_persona") diff --git a/backend/danswer/auth/api_key.py b/backend/danswer/auth/api_key.py new file mode 100644 index 00000000000..645130bef90 --- /dev/null +++ b/backend/danswer/auth/api_key.py @@ -0,0 +1,31 @@ +from fastapi import Depends +from fastapi import HTTPException +from fastapi import Request +from sqlalchemy import select +from sqlalchemy.orm import Session + +from danswer.db.engine import get_session +from danswer.db.models import ApiKey +from danswer.utils.logger import setup_logger + + +logger = setup_logger() + +_API_KEY_HEADER = "X-API-Key" + + +def validate_api_key(request: Request, db_session: Session = Depends(get_session)): + if _API_KEY_HEADER not in request.headers: + return None + + api_key_value = request.headers.get(_API_KEY_HEADER) + if not api_key_value: + raise HTTPException(status_code=401, detail="Missing API key") + + api_key = db_session.scalar( + select(ApiKey).where(ApiKey.hashed_api_key == api_key_value) + ) + if not api_key: + raise HTTPException(status_code=401, detail="Invalid API key") + + return None diff --git a/backend/danswer/chat/chat_utils.py b/backend/danswer/chat/chat_utils.py index f4b0b2e02c5..54fd072e172 100644 --- a/backend/danswer/chat/chat_utils.py +++ b/backend/danswer/chat/chat_utils.py @@ -114,27 +114,29 @@ def combine_message_chain( return "\n\n".join(message_strs) -def reorganize_citations( - answer: str, citations: list[CitationInfo] -) -> tuple[str, list[CitationInfo]]: - """For a complete, citation-aware response, we want to reorganize the citations so that +def reorganize_citations(answer: str, citations: list) -> tuple[str, list]: + """ + For a complete, citation-aware response, we want to reorganize the citations so that they are in the order of the documents that were used in the response. This just looks nicer / avoids - confusion ("Why is there [7] when only 2 documents are cited?").""" + confusion ("Why is there [7] when only 2 documents are cited?"). - # Regular expression to find all instances of [[x]](LINK) - pattern = r"\[\[(.*?)\]\]\((.*?)\)" + Now also handles citations in the format [number] in addition to [[number]](LINK). + """ + + pattern = r"\[\[(\d+)\]\]\((.*?)\)|\[(\d+)\]" all_citation_matches = re.findall(pattern, answer) new_citation_info: dict[int, CitationInfo] = {} for citation_match in all_citation_matches: try: - citation_num = int(citation_match[0]) + citation_str = citation_match[0] if citation_match[0] else citation_match[2] + citation_num = int(citation_str) if citation_num in new_citation_info: continue matching_citation = next( - iter([c for c in citations if c.citation_num == int(citation_num)]), + (c for c in citations if c.citation_num == citation_num), None, ) if matching_citation is None: @@ -149,16 +151,28 @@ def reorganize_citations( # Function to replace citations with their new number def slack_link_format(match: re.Match) -> str: - link_text = match.group(1) - try: - citation_num = int(link_text) - if citation_num in new_citation_info: - link_text = new_citation_info[citation_num].citation_num - except Exception: - pass - - link_url = match.group(2) - return f"[[{link_text}]]({link_url})" + # Case 1: Linked citation ([[number]](LINK)) + if match.group(1): + link_text = match.group(1) + try: + citation_num = int(link_text) + if citation_num in new_citation_info: + link_text = new_citation_info[citation_num].citation_num + except Exception: + pass + link_url = match.group(2) + return f"[[{link_text}]]({link_url})" + # Case 2: Non-linked citation ([number]) + elif match.group(3): + try: + citation_num = int(match.group(3)) + if citation_num in new_citation_info: + citation_num = new_citation_info[citation_num].citation_num + except Exception: + pass + return f"[{citation_num}]" + else: + return match.group(0) # Substitute all matches in the input text new_answer = re.sub(pattern, slack_link_format, answer) diff --git a/backend/danswer/configs/app_configs.py b/backend/danswer/configs/app_configs.py index f4e8ce14158..e447c189a47 100644 --- a/backend/danswer/configs/app_configs.py +++ b/backend/danswer/configs/app_configs.py @@ -198,6 +198,10 @@ for ignored_tag in os.environ.get("JIRA_CONNECTOR_LABELS_TO_SKIP", "").split(",") if ignored_tag ] +# Maximum size for Jira tickets in bytes (default: 100KB) +JIRA_CONNECTOR_MAX_TICKET_SIZE = int( + os.environ.get("JIRA_CONNECTOR_MAX_TICKET_SIZE", 100 * 1024) +) GONG_CONNECTOR_START_TIME = os.environ.get("GONG_CONNECTOR_START_TIME") diff --git a/backend/danswer/configs/constants.py b/backend/danswer/configs/constants.py index b29d3558b84..db445958e61 100644 --- a/backend/danswer/configs/constants.py +++ b/backend/danswer/configs/constants.py @@ -95,6 +95,7 @@ class DocumentSource(str, Enum): SHAREPOINT = "sharepoint" TEAMS = "teams" SALESFORCE = "salesforce" + SFKBARTICLES = "sfkbarticles" DISCOURSE = "discourse" AXERO = "axero" CLICKUP = "clickup" diff --git a/backend/danswer/connectors/danswer_jira/connector.py b/backend/danswer/connectors/danswer_jira/connector.py index da525146f9f..e70b28a1a2c 100644 --- a/backend/danswer/connectors/danswer_jira/connector.py +++ b/backend/danswer/connectors/danswer_jira/connector.py @@ -1,21 +1,25 @@ import os +from collections.abc import Iterable from datetime import datetime from datetime import timezone from typing import Any -from urllib.parse import urlparse from jira import JIRA from jira.resources import Issue from danswer.configs.app_configs import INDEX_BATCH_SIZE from danswer.configs.app_configs import JIRA_CONNECTOR_LABELS_TO_SKIP +from danswer.configs.app_configs import JIRA_CONNECTOR_MAX_TICKET_SIZE from danswer.configs.constants import DocumentSource from danswer.connectors.cross_connector_utils.miscellaneous_utils import time_str_to_utc +from danswer.connectors.danswer_jira.utils import best_effort_basic_expert_info +from danswer.connectors.danswer_jira.utils import best_effort_get_field_from_issue +from danswer.connectors.danswer_jira.utils import extract_text_from_content +from danswer.connectors.danswer_jira.utils import get_comment_strs from danswer.connectors.interfaces import GenerateDocumentsOutput from danswer.connectors.interfaces import LoadConnector from danswer.connectors.interfaces import PollConnector from danswer.connectors.interfaces import SecondsSinceUnixEpoch -from danswer.connectors.models import BasicExpertInfo from danswer.connectors.models import ConnectorMissingCredentialError from danswer.connectors.models import Document from danswer.connectors.models import Section @@ -23,174 +27,130 @@ logger = setup_logger() -PROJECT_URL_PAT = "projects" -JIRA_API_VERSION = os.environ.get("JIRA_API_VERSION") or "2" - - -def extract_jira_project(url: str) -> tuple[str, str]: - parsed_url = urlparse(url) - jira_base = parsed_url.scheme + "://" + parsed_url.netloc - - # Split the path by '/' and find the position of 'projects' to get the project name - split_path = parsed_url.path.split("/") - if PROJECT_URL_PAT in split_path: - project_pos = split_path.index(PROJECT_URL_PAT) - if len(split_path) > project_pos + 1: - jira_project = split_path[project_pos + 1] - else: - raise ValueError("No project name found in the URL") - else: - raise ValueError("'projects' not found in the URL") - - return jira_base, jira_project - -def extract_text_from_content(content: dict) -> str: - texts = [] - if "content" in content: - for block in content["content"]: - if "content" in block: - for item in block["content"]: - if item["type"] == "text": - texts.append(item["text"]) - return " ".join(texts) - - -def best_effort_get_field_from_issue(jira_issue: Issue, field: str) -> Any: - if hasattr(jira_issue.fields, field): - return getattr(jira_issue.fields, field) +JIRA_API_VERSION = os.environ.get("JIRA_API_VERSION") or "2" +_JIRA_FULL_PAGE_SIZE = 50 - try: - return jira_issue.raw["fields"][field] - except Exception: - return None +def _paginate_jql_search( + jira_client: JIRA, + jql: str, + max_results: int, + fields: str | None = None, +) -> Iterable[Issue]: + start = 0 + while True: + logger.debug( + f"Fetching Jira issues with JQL: {jql}, " + f"starting at {start}, max results: {max_results}" + ) + issues = jira_client.search_issues( + jql_str=jql, + startAt=start, + maxResults=max_results, + fields=fields, + ) -def _get_comment_strs( - jira: Issue, comment_email_blacklist: tuple[str, ...] = () -) -> list[str]: - comment_strs = [] - for comment in jira.fields.comment.comments: - try: - if hasattr(comment, "body"): - body_text = extract_text_from_content(comment.raw["body"]) - elif hasattr(comment, "raw"): - body = comment.raw.get("body", "No body content available") - body_text = ( - extract_text_from_content(body) if isinstance(body, dict) else body - ) + for issue in issues: + if isinstance(issue, Issue): + yield issue else: - body_text = "No body attribute found" - - if ( - hasattr(comment, "author") - and comment.author.emailAddress in comment_email_blacklist - ): - continue # Skip adding comment if author's email is in blacklist + raise Exception(f"Found Jira object not of type Issue: {issue}") - comment_strs.append(body_text) - except Exception as e: - logger.error(f"Failed to process comment due to an error: {e}") - continue + if len(issues) < max_results: + break - return comment_strs + start += max_results def fetch_jira_issues_batch( - jql: str, - start_index: int, jira_client: JIRA, - batch_size: int = INDEX_BATCH_SIZE, + jql: str, + batch_size: int, comment_email_blacklist: tuple[str, ...] = (), labels_to_skip: set[str] | None = None, -) -> tuple[list[Document], int]: - doc_batch = [] - - batch = jira_client.search_issues( - jql, - startAt=start_index, - maxResults=batch_size, - ) +) -> Iterable[Document]: + for issue in _paginate_jql_search( + jira_client=jira_client, + jql=jql, + max_results=batch_size, + ): + if labels_to_skip: + if any(label in issue.fields.labels for label in labels_to_skip): + logger.info( + f"Skipping {issue.key} because it has a label to skip. Found " + f"labels: {issue.fields.labels}. Labels to skip: {labels_to_skip}." + ) + continue - for jira in batch: - if type(jira) != Issue: - logger.warning(f"Found Jira object not of type Issue {jira}") - continue + description = ( + issue.fields.description + if JIRA_API_VERSION == "2" + else extract_text_from_content(issue.raw["fields"]["description"]) + ) + comments = get_comment_strs( + issue=issue, + comment_email_blacklist=comment_email_blacklist, + ) + ticket_content = f"{description}\n" + "\n".join( + [f"Comment: {comment}" for comment in comments if comment] + ) - if labels_to_skip and any( - label in jira.fields.labels for label in labels_to_skip - ): + # Check ticket size + if len(ticket_content.encode("utf-8")) > JIRA_CONNECTOR_MAX_TICKET_SIZE: logger.info( - f"Skipping {jira.key} because it has a label to skip. Found " - f"labels: {jira.fields.labels}. Labels to skip: {labels_to_skip}." + f"Skipping {issue.key} because it exceeds the maximum size of " + f"{JIRA_CONNECTOR_MAX_TICKET_SIZE} bytes." ) continue - comments = _get_comment_strs(jira, comment_email_blacklist) - semantic_rep = ( - f"{jira.fields.description}\n" - if jira.fields.description - else "" + "\n".join([f"Comment: {comment}" for comment in comments]) - ) - - page_url = f"{jira_client.client_info()}/browse/{jira.key}" + page_url = f"{jira_client.client_info()}/browse/{issue.key}" people = set() try: - people.add( - BasicExpertInfo( - display_name=jira.fields.creator.displayName, - email=jira.fields.creator.emailAddress, - ) - ) + creator = best_effort_get_field_from_issue(issue, "creator") + if basic_expert_info := best_effort_basic_expert_info(creator): + people.add(basic_expert_info) except Exception: # Author should exist but if not, doesn't matter pass try: - people.add( - BasicExpertInfo( - display_name=jira.fields.assignee.displayName, # type: ignore - email=jira.fields.assignee.emailAddress, # type: ignore - ) - ) + assignee = best_effort_get_field_from_issue(issue, "assignee") + if basic_expert_info := best_effort_basic_expert_info(assignee): + people.add(basic_expert_info) except Exception: # Author should exist but if not, doesn't matter pass metadata_dict = {} - priority = best_effort_get_field_from_issue(jira, "priority") - if priority: + if priority := best_effort_get_field_from_issue(issue, "priority"): metadata_dict["priority"] = priority.name - status = best_effort_get_field_from_issue(jira, "status") - if status: + if status := best_effort_get_field_from_issue(issue, "status"): metadata_dict["status"] = status.name - resolution = best_effort_get_field_from_issue(jira, "resolution") - if resolution: + if resolution := best_effort_get_field_from_issue(issue, "resolution"): metadata_dict["resolution"] = resolution.name - labels = best_effort_get_field_from_issue(jira, "labels") - if labels: + if labels := best_effort_get_field_from_issue(issue, "labels"): metadata_dict["label"] = labels - doc_batch.append( - Document( - id=page_url, - sections=[Section(link=page_url, text=semantic_rep)], - source=DocumentSource.JIRA, - semantic_identifier=jira.fields.summary, - doc_updated_at=time_str_to_utc(jira.fields.updated), - primary_owners=list(people) or None, - # TODO add secondary_owners (commenters) if needed - metadata=metadata_dict, - ) + yield Document( + id=page_url, + sections=[Section(link=page_url, text=ticket_content)], + source=DocumentSource.JIRA, + semantic_identifier=f"{issue.key}: {issue.fields.summary}", + title=f"{issue.key} {issue.fields.summary}", + doc_updated_at=time_str_to_utc(issue.fields.updated), + primary_owners=list(people) or None, + # TODO add secondary_owners (commenters) if needed + metadata=metadata_dict, ) - return doc_batch, len(batch) class JiraConnector(LoadConnector, PollConnector): def __init__( self, - jira_project_url: str, + jira_base_url: str, + jira_filter: str, comment_email_blacklist: list[str] | None = None, batch_size: int = INDEX_BATCH_SIZE, # if a ticket has one of the labels specified in this list, we will just @@ -199,28 +159,35 @@ def __init__( labels_to_skip: list[str] = JIRA_CONNECTOR_LABELS_TO_SKIP, ) -> None: self.batch_size = batch_size - self.jira_base, self.jira_project = extract_jira_project(jira_project_url) - self.jira_client: JIRA | None = None + self.jira_base = jira_base_url + self._jira_client: JIRA | None = None self._comment_email_blacklist = comment_email_blacklist or [] self.labels_to_skip = set(labels_to_skip) + self.jira_filter = jira_filter @property def comment_email_blacklist(self) -> tuple: return tuple(email.strip() for email in self._comment_email_blacklist) + @property + def jira_client(self) -> JIRA: + if self._jira_client is None: + raise ConnectorMissingCredentialError("Jira") + return self._jira_client + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: api_token = credentials["jira_api_token"] # if user provide an email we assume it's cloud if "jira_user_email" in credentials: email = credentials["jira_user_email"] - self.jira_client = JIRA( + self._jira_client = JIRA( basic_auth=(email, api_token), server=self.jira_base, options={"rest_api_version": JIRA_API_VERSION}, ) else: - self.jira_client = JIRA( + self._jira_client = JIRA( token_auth=api_token, server=self.jira_base, options={"rest_api_version": JIRA_API_VERSION}, @@ -228,26 +195,22 @@ def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None return None def load_from_state(self) -> GenerateDocumentsOutput: - if self.jira_client is None: - raise ConnectorMissingCredentialError("Jira") - - start_ind = 0 - while True: - doc_batch, fetched_batch_size = fetch_jira_issues_batch( - jql=f"project = {self.jira_project}", - start_index=start_ind, - jira_client=self.jira_client, - batch_size=self.batch_size, - comment_email_blacklist=self.comment_email_blacklist, - labels_to_skip=self.labels_to_skip, - ) - - if doc_batch: - yield doc_batch + jql = f"project = {self.quoted_jira_project}" + + document_batch = [] + for doc in fetch_jira_issues_batch( + jira_client=self.jira_client, + jql=jql, + batch_size=_JIRA_FULL_PAGE_SIZE, + comment_email_blacklist=self.comment_email_blacklist, + labels_to_skip=self.labels_to_skip, + ): + document_batch.append(doc) + if len(document_batch) >= self.batch_size: + yield document_batch + document_batch = [] - start_ind += fetched_batch_size - if fetched_batch_size < self.batch_size: - break + yield document_batch def poll_source( self, start: SecondsSinceUnixEpoch, end: SecondsSinceUnixEpoch @@ -263,36 +226,31 @@ def poll_source( ) jql = ( - f"project = {self.jira_project} AND " + f"{self.jira_filter} AND " f"updated >= '{start_date_str}' AND " f"updated <= '{end_date_str}'" ) - start_ind = 0 - while True: - doc_batch, fetched_batch_size = fetch_jira_issues_batch( - jql=jql, - start_index=start_ind, - jira_client=self.jira_client, - batch_size=self.batch_size, - comment_email_blacklist=self.comment_email_blacklist, - labels_to_skip=self.labels_to_skip, - ) - - if doc_batch: - yield doc_batch + document_batch = [] + for doc in fetch_jira_issues_batch( + jira_client=self.jira_client, + jql=jql, + batch_size=_JIRA_FULL_PAGE_SIZE, + comment_email_blacklist=self.comment_email_blacklist, + labels_to_skip=self.labels_to_skip, + ): + document_batch.append(doc) + if len(document_batch) >= self.batch_size: + yield document_batch + document_batch = [] - start_ind += fetched_batch_size - if fetched_batch_size < self.batch_size: - break + yield document_batch if __name__ == "__main__": import os - connector = JiraConnector( - os.environ["JIRA_PROJECT_URL"], comment_email_blacklist=[] - ) + connector = JiraConnector(os.environ["JIRA_FILTERS"], comment_email_blacklist=[]) connector.load_credentials( { "jira_user_email": os.environ["JIRA_USER_EMAIL"], @@ -300,4 +258,4 @@ def poll_source( } ) document_batches = connector.load_from_state() - print(next(document_batches)) \ No newline at end of file + print(next(document_batches)) diff --git a/backend/danswer/connectors/danswer_jira/utils.py b/backend/danswer/connectors/danswer_jira/utils.py index 506f5eff75e..b8c62bf5780 100644 --- a/backend/danswer/connectors/danswer_jira/utils.py +++ b/backend/danswer/connectors/danswer_jira/utils.py @@ -1,4 +1,5 @@ """Module with custom fields processing functions""" +import os from typing import Any from typing import List @@ -7,10 +8,80 @@ from jira.resources import Issue from jira.resources import User +from danswer.connectors.models import BasicExpertInfo from danswer.utils.logger import setup_logger logger = setup_logger() +JIRA_API_VERSION = os.environ.get("JIRA_API_VERSION") or "2" + + +def best_effort_basic_expert_info(obj: Any) -> BasicExpertInfo | None: + display_name = None + email = None + if hasattr(obj, "display_name"): + display_name = obj.display_name + else: + display_name = obj.get("displayName") + + if hasattr(obj, "emailAddress"): + email = obj.emailAddress + else: + email = obj.get("emailAddress") + + if not email and not display_name: + return None + + return BasicExpertInfo(display_name=display_name, email=email) + + +def best_effort_get_field_from_issue(jira_issue: Issue, field: str) -> Any: + if hasattr(jira_issue.fields, field): + return getattr(jira_issue.fields, field) + + try: + return jira_issue.raw["fields"][field] + except Exception: + return None + + +def extract_text_from_content(content: dict) -> str: + texts = [] + if "content" in content: + for block in content["content"]: + if "content" in block: + for item in block["content"]: + if item["type"] == "text": + texts.append(item["text"]) + return " ".join(texts) + + +def get_comment_strs( + issue: Issue, comment_email_blacklist: tuple[str, ...] = () +) -> list[str]: + comment_strs = [] + for comment in issue.fields.comment.comments: + try: + body_text = ( + comment.body + if JIRA_API_VERSION == "2" + else extract_text_from_content(comment.raw["body"]) + ) + + if ( + hasattr(comment, "author") + and hasattr(comment.author, "emailAddress") + and comment.author.emailAddress in comment_email_blacklist + ): + continue # Skip adding comment if author's email is in blacklist + + comment_strs.append(body_text) + except Exception as e: + logger.error(f"Failed to process comment due to an error: {e}") + continue + + return comment_strs + class CustomFieldExtractor: @staticmethod diff --git a/backend/danswer/connectors/factory.py b/backend/danswer/connectors/factory.py index 1a3d605d3a5..2f885bf2b7d 100644 --- a/backend/danswer/connectors/factory.py +++ b/backend/danswer/connectors/factory.py @@ -34,6 +34,7 @@ from danswer.connectors.productboard.connector import ProductboardConnector from danswer.connectors.requesttracker.connector import RequestTrackerConnector from danswer.connectors.salesforce.connector import SalesforceConnector +from danswer.connectors.sfkbarticles.connector import SfKbArticlesConnector from danswer.connectors.sharepoint.connector import SharepointConnector from danswer.connectors.slab.connector import SlabConnector from danswer.connectors.slack.connector import SlackPollConnector @@ -86,6 +87,7 @@ def identify_connector_class( DocumentSource.SHAREPOINT: SharepointConnector, DocumentSource.TEAMS: TeamsConnector, DocumentSource.SALESFORCE: SalesforceConnector, + DocumentSource.SFKBARTICLES: SfKbArticlesConnector, DocumentSource.DISCOURSE: DiscourseConnector, DocumentSource.AXERO: AxeroConnector, DocumentSource.CLICKUP: ClickupConnector, diff --git a/backend/danswer/connectors/sfkbarticles/__init__.py b/backend/danswer/connectors/sfkbarticles/__init__.py new file mode 100644 index 00000000000..e69de29bb2d diff --git a/backend/danswer/connectors/sfkbarticles/connector.py b/backend/danswer/connectors/sfkbarticles/connector.py new file mode 100644 index 00000000000..91bed54617c --- /dev/null +++ b/backend/danswer/connectors/sfkbarticles/connector.py @@ -0,0 +1,256 @@ +import os +from datetime import datetime +from datetime import timezone +from typing import Any +from typing import Tuple + +import requests + +from danswer.configs.app_configs import INDEX_BATCH_SIZE +from danswer.configs.constants import DocumentSource +from danswer.connectors.cross_connector_utils.miscellaneous_utils import time_str_to_utc +from danswer.connectors.interfaces import GenerateDocumentsOutput +from danswer.connectors.interfaces import LoadConnector +from danswer.connectors.interfaces import PollConnector +from danswer.connectors.interfaces import SecondsSinceUnixEpoch +from danswer.connectors.models import BasicExpertInfo +from danswer.connectors.models import Document +from danswer.connectors.models import Section +from danswer.connectors.sfkbarticles.utils import extract_dict_text +from danswer.utils.logger import setup_logger + +ID_PREFIX = "SALESFORCE_" +AUTH_URL = "https://login.salesforce.com/services/oauth2/token" + +logger = setup_logger() + + +class SfKbArticlesConnector(LoadConnector, PollConnector): + def __init__( + self, + batch_size: int = INDEX_BATCH_SIZE, + requested_objects: list[str] = [], + ) -> None: + self.batch_size = batch_size + self.product_component_list = ( + [obj.strip() for obj in requested_objects[0].split(",")] + if requested_objects + else None + ) + + def load_credentials(self, credentials: dict[str, Any]) -> dict[str, Any] | None: + self.client_id = credentials["sf_client_id"] + self.client_secret = credentials["sf_client_secret"] + self.username = credentials["sf_username"] + self.password = credentials["sf_password"] + + self.access_token, self.instance_url = self._get_access_token() + + self.headers = { + "Authorization": f"Bearer {self.access_token}", + "Content-Type": "application/json", + } + + def _get_access_token(self) -> Tuple[str, str]: + """ + Authenticates with Salesforce and retrieves the access token & instance URL. + """ + payload = { + "grant_type": "password", + "client_id": self.client_id, + "client_secret": self.client_secret, + "username": self.username, + "password": f"{self.password}", + } + response = requests.post(AUTH_URL, data=payload) + if response.status_code != 200: + logger.error(f"Authentication failed: {response.text}") + raise Exception("Failed to authenticate with Salesforce.") + + data = response.json() + logger.info("Successfully authenticated with Salesforce.") + return data["access_token"], data["instance_url"] + + def _convert_object_instance_to_document( + self, object_dict: dict[str, Any] + ) -> Document: + salesforce_id = object_dict["Id"] + danswer_salesforce_id = f"{ID_PREFIX}{salesforce_id}" + extracted_link = f"{self.instance_url}/{salesforce_id}" + extracted_doc_updated_at = time_str_to_utc(object_dict["LastModifiedDate"]) + extracted_object_text = extract_dict_text(object_dict) + extracted_semantic_identifier = object_dict.get("Title", "Unknown Object") + extracted_primary_owners = [ + BasicExpertInfo( + display_name=self._get_name_from_id(object_dict["LastModifiedById"]) + ) + ] + + doc = Document( + id=danswer_salesforce_id, + sections=[Section(link=extracted_link, text=extracted_object_text)], + source=DocumentSource.SFKBARTICLES, + semantic_identifier=extracted_semantic_identifier, + doc_updated_at=extracted_doc_updated_at, + primary_owners=extracted_primary_owners, + metadata={}, + ) + return doc + + def _get_name_from_id(self, id: str) -> str: + """ + Fetches the name of a Salesforce user based on their ID. + """ + query = f"SELECT Name FROM User WHERE Id = '{id}'" + url = f"{self.instance_url}/services/data/v56.0/query" + params = {"q": query} + + response = requests.get(url, headers=self.headers, params=params) + if response.status_code != 200: + logger.error(f"Failed to fetch name for ID {id}: {response.text}") + return "Unknown" + + data = response.json() + records = data.get("records", []) + if not records: + logger.warning(f"No name found for ID {id}") + return "Unknown" + + return records[0].get("Name", "Unknown") + + def build_salesforce_query( + self, + parent_objects: list[str], + start: datetime | None = None, + end: datetime | None = None, + ) -> str: + """ + Builds the Salesforce query dynamically. + If product_component_list is empty, it fetches all Product_Component__c values. + Otherwise, it filters using the provided product_component_list. + """ + if not parent_objects: + product_filter = "" # No filter, fetch all + else: + product_components_str = ", ".join( + [ + f"'{component.strip()}'" + for component in parent_objects + if component.strip() + ] + ) + product_filter = f"AND Product_Component__c IN ({product_components_str})" + + query = f""" + SELECT + Id, + Title, + Summary, + Product_Component__c, + Product_Component_Version__c, + Question_Problem__c, + Resolution__c, + Sub_Component__c, + IsVisibleInPkb, + IsVisibleInCsp, + IsVisibleInPrm, + ArticleCreatedDate, + ArticleNumber, + CreatedDate, + ArticleTotalViewCount, + FirstPublishedDate, + IsDeleted, + IsLatestVersion, + Language, + LastModifiedDate, + LastModifiedById, + LastPublishedDate, + Orchestrator_Version__c, + Studio_Version__c + FROM Knowledge__kav + WHERE Language='en_US' + AND PublishStatus = 'Online' + AND IsDeleted = FALSE + {product_filter} + """.strip() + + if start: + start_str = start.strftime("%Y-%m-%dT%H:%M:%S.000Z") + query += f" AND LastModifiedDate >= {start_str}" + if end: + end_str = end.strftime("%Y-%m-%dT%H:%M:%S.000Z") + query += f" AND LastModifiedDate <= {end_str}" + + return query + + def _fetch_from_salesforce( + self, + start: datetime | None = None, + end: datetime | None = None, + ) -> GenerateDocumentsOutput: + query = self.build_salesforce_query(self.product_component_list, start, end) + doc_batch: list[Document] = [] + query_results: dict = {} + + url = f"{self.instance_url}/services/data/v56.0/query" + params = {"q": query} + query_result = requests.get( + url, headers=self.headers, params=params if "q" in params else None + ) + + while url: + query_result = requests.get( + url, headers=self.headers, params=params if "q" in params else None + ) + data = query_result.json() + + if isinstance(data, list): + error_message = "; ".join( + error.get("message", "Unknown error") for error in data + ) + raise Exception(f"Salesforce API error: {error_message}") + + if "records" in data: + for record_dict in data["records"]: + query_results.setdefault(record_dict["Id"], {}).update(record_dict) + + url = data.get("nextRecordsUrl", None) + if url: + url = f"{self.instance_url}{url}" + + for combined_object_dict in query_results.values(): + doc_batch.append( + self._convert_object_instance_to_document(combined_object_dict) + ) + if len(doc_batch) > self.batch_size: + yield doc_batch + doc_batch = [] + + yield doc_batch + + def load_from_state(self) -> GenerateDocumentsOutput: + return self._fetch_from_salesforce() + + def poll_source( + self, start: SecondsSinceUnixEpoch, end: SecondsSinceUnixEpoch + ) -> GenerateDocumentsOutput: + start_datetime = datetime.fromtimestamp(start, tz=timezone.utc) + end_datetime = datetime.fromtimestamp(end, tz=timezone.utc) + return self._fetch_from_salesforce(start=start_datetime, end=end_datetime) + + +if __name__ == "__main__": + connector = SfKbArticlesConnector( + requested_objects=os.environ["REQUESTED_OBJECTS"].split(",") + ) + + connector.load_credentials( + { + "sf_client_id": os.environ["SF_CLIENT_ID"], + "sf_client_secret": os.environ["SF_CLIENT_SECRET"], + "sf_username": os.environ["SF_USERNAME"], + "sf_password": os.environ["SF_PASSWORD"], + } + ) + document_batches = connector.load_from_state() + print(next(document_batches)) diff --git a/backend/danswer/connectors/sfkbarticles/utils.py b/backend/danswer/connectors/sfkbarticles/utils.py new file mode 100644 index 00000000000..15aa5512cd0 --- /dev/null +++ b/backend/danswer/connectors/sfkbarticles/utils.py @@ -0,0 +1,95 @@ +import re +from typing import Union + +from bs4 import BeautifulSoup + + +SF_JSON_FILTER = r"Id$|Date$|stamp$|url$" + + +def _clean_salesforce_dict(data: Union[dict, list]) -> Union[dict, list]: + if isinstance(data, dict): + if "records" in data.keys(): + data = data["records"] + if isinstance(data, dict): + if "attributes" in data.keys(): + if isinstance(data["attributes"], dict): + data.update(data.pop("attributes")) + + if isinstance(data, dict): + filtered_dict = {} + for key, value in data.items(): + if not re.search(SF_JSON_FILTER, key, re.IGNORECASE): + if "__c" in key: # remove the custom object indicator for display + key = key[:-3] + if isinstance(value, (dict, list)): + filtered_value = _clean_salesforce_dict(value) + if filtered_value: + filtered_dict[key] = filtered_value + elif value is not None: + filtered_dict[key] = value + return filtered_dict + elif isinstance(data, list): + filtered_list = [] + for item in data: + if isinstance(item, (dict, list)): + filtered_item = _clean_salesforce_dict(item) + if filtered_item: + filtered_list.append(filtered_item) + elif item is not None: + filtered_list.append(filtered_item) + return filtered_list + else: + return data + + +def _json_to_natural_language(data: Union[dict, list], indent: int = 0) -> str: + result = [] + indent_str = " " * indent + + if isinstance(data, dict): + for key, value in data.items(): + if isinstance(value, (dict, list)): + result.append(f"{indent_str}{key}:") + result.append(_json_to_natural_language(value, indent + 2)) + else: + result.append(f"{indent_str}{key}: {value}") + elif isinstance(data, list): + for item in data: + result.append(_json_to_natural_language(item, indent)) + else: + result.append(f"{indent_str}{data}") + + return "\n".join(result) + + +def extract_dict_text(raw_dict: dict) -> str: + processed_dict = _clean_salesforce_dict(raw_dict) + + if "Resolution" in processed_dict: + processed_dict["Resolution"] = clean_html(processed_dict["Resolution"]) + + natural_language_dict = _json_to_natural_language(processed_dict) + return natural_language_dict + + +def clean_html(html_content: str) -> str: + """ + Cleans the HTML content by removing tags and returning the plain text, + while preserving the links. + :param html_content: HTML content as string + :return: Cleaned text with preserved links + """ + soup = BeautifulSoup(html_content, "lxml") + + # Replace tags with their text and the href as a link in brackets + for a_tag in soup.find_all("a", href=True): + a_tag.insert_before(f"[{a_tag.get_text()}]({a_tag['href']})") + a_tag.decompose() + + cleaned_text = soup.get_text(separator="\n", strip=True) + + # Replace non-breaking spaces (\xa0) with regular spaces + cleaned_text = cleaned_text.replace("\xa0", " ") + + return cleaned_text diff --git a/backend/danswer/connectors/slack/utils.py b/backend/danswer/connectors/slack/utils.py index 21bae6571d8..48a591f2aca 100644 --- a/backend/danswer/connectors/slack/utils.py +++ b/backend/danswer/connectors/slack/utils.py @@ -93,13 +93,13 @@ def rate_limited_call(**kwargs: Any) -> SlackResponse: error = "unknown error" if error == "ratelimited": - # Handle rate limiting: get the 'Retry-After' header value and sleep for that duration - retry_after = int(e.response.headers['Retry-After']) - logger.info( - f"Slack call rate limited, retrying after {retry_after} seconds. Exception: {e}" - ) - time.sleep(retry_after) - continue + # Handle rate limiting: get the 'Retry-After' header value and sleep for that duration + retry_after = int(e.response.headers["Retry-After"]) + logger.info( + f"Slack call rate limited, retrying after {retry_after} seconds. Exception: {e}" + ) + time.sleep(retry_after) + continue elif error in ["already_reacted", "no_reaction"]: logger.info("not here already_reacted") # The response isn't used for reactions, this is basically just a pass @@ -278,3 +278,9 @@ def replace_special_catchall(message: str) -> str: def add_zero_width_whitespace_after_tag(message: str) -> str: """Add a 0 width whitespace after every @""" return message.replace("@", "@\u200B") + + @staticmethod + def handle_bold_syntax_for_slack(text: str) -> str: + """Replace instances of '**' with a single '*'""" + corrected_text = text.replace("**", "*") + return corrected_text diff --git a/backend/danswer/danswerbot/slack/handlers/handle_buttons.py b/backend/danswer/danswerbot/slack/handlers/handle_buttons.py index 3a0209b076f..bddb40ae829 100644 --- a/backend/danswer/danswerbot/slack/handlers/handle_buttons.py +++ b/backend/danswer/danswerbot/slack/handlers/handle_buttons.py @@ -7,6 +7,7 @@ from slack_sdk.socket_mode import SocketModeClient from slack_sdk.socket_mode.request import SocketModeRequest from sqlalchemy.orm import Session +from sqlalchemy.orm.exc import NoResultFound from danswer.configs.constants import SearchFeedbackType from danswer.configs.danswerbot_configs import DANSWER_FOLLOWUP_EMOJI @@ -32,6 +33,10 @@ from danswer.db.engine import get_sqlalchemy_engine from danswer.db.feedback import create_chat_message_feedback from danswer.db.feedback import create_doc_retrieval_feedback +from danswer.db.persona import fetch_persona_by_id +from danswer.db.users import add_slack_persona_for_user +from danswer.db.users import add_user_slack_persona +from danswer.db.users import fetch_user_slack_persona from danswer.document_index.document_index_utils import get_both_index_names from danswer.document_index.factory import get_default_document_index from danswer.utils.logger import setup_logger @@ -293,3 +298,63 @@ def handle_followup_resolved_button( thread_ts=thread_ts, unfurl=False, ) + + +def handle_persona_selection(req: SocketModeRequest, client: SocketModeClient) -> None: + action = cast(dict[str, Any], req.payload.get("actions", [])[0]) + user_id = req.payload["user"]["id"] + channel_id = req.payload["container"]["channel_id"] + persona_id = action.get("value") + message_ts_to_respond_to = req.payload.get("container", {}).get("thread_ts") + + with Session(get_sqlalchemy_engine()) as db_session: + try: + persona = fetch_persona_by_id(db_session=db_session, persona_id=persona_id) + + if persona is None: + respond_in_thread( + client=client.web_client, + channel=channel_id, + text="Persona not found.", + thread_ts=message_ts_to_respond_to, + ) + return + + user_slack_persona = fetch_user_slack_persona( + db_session=db_session, sender_id=user_id + ) + if user_slack_persona: + add_slack_persona_for_user( + db_session=db_session, + persona=persona, + user_slack_persona=user_slack_persona, + ) + response_text = f"Persona '{persona.name}' has been set!\n" + respond_in_thread( + client=client.web_client, + channel=channel_id, + text=response_text, + thread_ts=message_ts_to_respond_to, + ) + return + + else: + add_user_slack_persona( + db_session=db_session, sender_id=user_id, persona=persona + ) + respond_in_thread( + client=client.web_client, + channel=channel_id, + text=f"'{persona.name}' has been successfully set as the current persona.", + thread_ts=message_ts_to_respond_to, + ) + return + + except NoResultFound: + respond_in_thread( + client=client.web_client, + channel=channel_id, + text="Error in fetching persona", + thread_ts=message_ts_to_respond_to, + ) + return diff --git a/backend/danswer/danswerbot/slack/handlers/handle_message.py b/backend/danswer/danswerbot/slack/handlers/handle_message.py index 85ddfe70a6c..1e9d14b8869 100644 --- a/backend/danswer/danswerbot/slack/handlers/handle_message.py +++ b/backend/danswer/danswerbot/slack/handlers/handle_message.py @@ -1,6 +1,7 @@ import datetime import functools import logging +import re from collections.abc import Callable from typing import Any from typing import cast @@ -23,7 +24,6 @@ from danswer.configs.danswerbot_configs import DANSWER_BOT_NUM_RETRIES from danswer.configs.danswerbot_configs import DANSWER_BOT_TARGET_CHUNK_PERCENTAGE from danswer.configs.danswerbot_configs import DANSWER_BOT_USE_QUOTES -from danswer.configs.danswerbot_configs import DANSWER_FOLLOWUP_EMOJI from danswer.configs.danswerbot_configs import DANSWER_REACT_EMOJI from danswer.configs.danswerbot_configs import DISABLE_DANSWER_BOT_FILTER_DETECT from danswer.configs.danswerbot_configs import ENABLE_DANSWERBOT_REFLEXION @@ -47,6 +47,9 @@ from danswer.db.models import SlackBotConfig from danswer.db.models import SlackBotResponseType from danswer.db.persona import fetch_persona_by_id +from danswer.db.persona import get_persona_with_docset_and_prompts +from danswer.db.persona import get_personas +from danswer.db.users import fetch_user_slack_persona from danswer.llm.answering.prompts.citations_prompt import ( compute_max_document_tokens_for_persona, ) @@ -67,6 +70,7 @@ srl = SlackRateLimiter() RT = TypeVar("RT") # return type +MAX_BUTTONS_PER_BLOCK = 25 # Slack limit def rate_limits( @@ -172,11 +176,23 @@ def remove_scheduled_feedback_reminder( ) +def contains_questionmark_outside_links(message: str) -> bool: + """ + Checks if the message contains a question mark outside of URLs. + """ + url_pattern = r"]+>|https?://\S+" + + message_without_links = re.sub(url_pattern, "", message) + + return "?" in message_without_links + + def handle_message( message_info: SlackMessageInfo, channel_config: SlackBotConfig | None, client: WebClient, feedback_reminder_id: str | None, + channel_name: str | None, num_retries: int = DANSWER_BOT_NUM_RETRIES, answer_generation_timeout: int = DANSWER_BOT_ANSWER_GENERATION_TIMEOUT, should_respond_with_error_msgs: bool = DANSWER_BOT_DISPLAY_ERROR_MSGS, @@ -206,9 +222,98 @@ def handle_message( bypass_filters = message_info.bypass_filters is_bot_msg = message_info.is_bot_msg is_bot_dm = message_info.is_bot_dm + persona_name = None + + if channel_name is None: + with Session(get_sqlalchemy_engine()) as db_session: + user_slack_persona = fetch_user_slack_persona( + db_session=db_session, sender_id=sender_id + ) + if user_slack_persona: + slack_persona_id = user_slack_persona.persona_id or None + persona = get_persona_with_docset_and_prompts( + persona_id=slack_persona_id, db_session=db_session + ) + persona_name = persona.name + else: + persona = None + else: + persona = channel_config.persona if channel_config else None + + if is_bot_msg and channel_name is None: + command = message_info.command + if command == "/personas": + with Session(get_sqlalchemy_engine()) as db_session: + personas = get_personas( + user_id=None, db_session=db_session, include_default=False + ) + + if not personas: + respond_in_thread( + client=client, + channel=channel, + text="No personas are available.", + thread_ts=message_ts_to_respond_to, + ) + return + + persona_buttons = [ + { + "type": "button", + "text": {"type": "plain_text", "text": persona.name}, + "value": str(persona.id), + "action_id": f"set_persona_{persona.id}", + } + for persona in personas + ] + + blocks = [ + { + "type": "section", + "text": { + "type": "mrkdwn", + "text": "Here are the available personas. Click on one to set it:", + }, + } + ] + + for i in range(0, len(persona_buttons), MAX_BUTTONS_PER_BLOCK): + blocks.append( + { + "type": "actions", + "elements": persona_buttons[i : i + MAX_BUTTONS_PER_BLOCK], + } + ) + + respond_in_thread( + client=client, + channel=channel, + blocks=blocks, + thread_ts=message_ts_to_respond_to, + ) + return + elif command == "/current_persona": + if persona_name is None: + respond_in_thread( + client=client, + channel=channel, + text="No persona is set. Please use the /personas command to set up a persona.", + thread_ts=message_ts_to_respond_to, + ) + return + else: + respond_in_thread( + client=client, + channel=channel, + text=f"Current persona : {persona_name}", + thread_ts=message_ts_to_respond_to, + ) + return + elif is_bot_msg and (command == "/personas" or command == "/current_persona"): + logger.info("The slash command was used in a channel, won't work") + return document_set_names: list[str] | None = None - persona = channel_config.persona if channel_config else None prompt = None if persona: document_set_names = [ @@ -247,10 +352,9 @@ def handle_message( if not bypass_filters and "answer_filters" in channel_conf: reflexion = "well_answered_postfilter" in channel_conf["answer_filters"] - if ( - "questionmark_prefilter" in channel_conf["answer_filters"] - and "?" not in messages[-1].message - ): + if "questionmark_prefilter" in channel_conf[ + "answer_filters" + ] and not contains_questionmark_outside_links(messages[-1].message): logger.info( "Skipping message since it does not contain a question mark" ) @@ -487,21 +591,22 @@ def _get_answer(new_message_request: DirectQARequest) -> OneShotQAResponse | Non except SlackApiError as e: logger.error(f"Failed to remove Reaction due to: {e}") - if answer.answer_valid is False: - logger.info( - "Answer was evaluated to be invalid, throwing it away without responding." - ) - update_emote_react( - emoji=DANSWER_FOLLOWUP_EMOJI, - channel=message_info.channel_to_respond, - message_ts=message_info.msg_to_respond, - remove=False, - client=client, - ) - - if answer.answer: - logger.debug(answer.answer) - return True + # Removing this as we are handling this with citations logic + # if answer.answer_valid is False: + # logger.info( + # "Answer was evaluated to be invalid, throwing it away without responding." + # ) + # update_emote_react( + # emoji=DANSWER_FOLLOWUP_EMOJI, + # channel=message_info.channel_to_respond, + # message_ts=message_info.msg_to_respond, + # remove=False, + # client=client, + # ) + + # if answer.answer: + # logger.debug(answer.answer) + # return True retrieval_info = answer.docs if not retrieval_info: @@ -571,6 +676,18 @@ def _get_answer(new_message_request: DirectQARequest) -> OneShotQAResponse | Non if matching_doc: cited_docs.append((citation.citation_num, matching_doc)) + if not cited_docs: + logger.info("Skipping response: No context documents cited for this query.") + update_emote_react( + emoji="no-idea", + channel=channel, + message_ts=message_ts_to_respond_to, + remove=False, + client=client, + ) + + return True + cited_docs.sort() citations_block = build_sources_blocks(cited_documents=cited_docs) elif priority_ordered_docs: diff --git a/backend/danswer/danswerbot/slack/listener.py b/backend/danswer/danswerbot/slack/listener.py index f8dfb600211..582ed2aa774 100644 --- a/backend/danswer/danswerbot/slack/listener.py +++ b/backend/danswer/danswerbot/slack/listener.py @@ -1,3 +1,4 @@ +import re import time from threading import Event from typing import Any @@ -27,6 +28,7 @@ from danswer.danswerbot.slack.handlers.handle_buttons import ( handle_followup_resolved_button, ) +from danswer.danswerbot.slack.handlers.handle_buttons import handle_persona_selection from danswer.danswerbot.slack.handlers.handle_buttons import handle_slack_feedback from danswer.danswerbot.slack.handlers.handle_message import handle_message from danswer.danswerbot.slack.handlers.handle_message import ( @@ -94,6 +96,16 @@ def prefilter_requests(req: SocketModeRequest, client: SocketModeClient) -> bool channel_specific_logger.error("Cannot respond to empty message - skipping") return False + if re.search(r"!darwin", msg, re.IGNORECASE): + channel_specific_logger.info("Ignoring message containing '!darwin'") + return False + + if re.search(r":announcement\d*:", msg): + channel_specific_logger.info( + "Ignoring message: contains the announcement emoji" + ) + return False + if ( req.payload.setdefault("event", {}).get("user", "") == _OFFICIAL_SLACKBOT_USER_ID @@ -167,7 +179,7 @@ def prefilter_requests(req: SocketModeRequest, client: SocketModeClient) -> bool and message_ts != thread_ts and event_type != "app_mention" and event.get("channel_type") != "im" - and event.get("subtype") != "thread_broadcast" + and event.get("subtype") != "thread_broadcast" ): channel_specific_logger.debug( "Skipping message since it is not the root of a thread" @@ -198,6 +210,21 @@ def prefilter_requests(req: SocketModeRequest, client: SocketModeClient) -> bool ) return False + # Do not respond to messages if the channel is tagged + payload = req.payload + event = payload.get("event", {}) + blocks = event.get("blocks", []) + for block in blocks: + if block.get("type") == "rich_text": + for element in block.get("elements", []): + if element.get("type") == "rich_text_section": + for sub_element in element.get("elements", []): + if sub_element.get("type") == "broadcast" and sub_element.get( + "range" + ) in {"channel", "here"}: + logger.info("Broadcast message detected; skipping reply.") + return False + logger.debug(f"Handling Slack request with Payload: '{req.payload}'") return True @@ -277,6 +304,7 @@ def build_request_details( channel = req.payload["channel_id"] msg = req.payload["text"] sender = req.payload["user_id"] + command = req.payload["command"] single_msg = ThreadMessage(message=msg, sender=None, role=MessageType.USER) @@ -288,6 +316,7 @@ def build_request_details( bypass_filters=True, is_bot_msg=True, is_bot_dm=False, + command=command, ) raise RuntimeError("Programming fault, this should never happen.") @@ -356,6 +385,7 @@ def process_message( channel_config=slack_bot_config, client=client.web_client, feedback_reminder_id=feedback_reminder_id, + channel_name=channel_name, ) if failed: @@ -391,6 +421,9 @@ def action_routing(req: SocketModeRequest, client: SocketModeClient) -> None: return handle_followup_resolved_button(req, client, immediate=True) elif action["action_id"] == FOLLOWUP_BUTTON_RESOLVED_ACTION_ID: return handle_followup_resolved_button(req, client, immediate=False) + elif action["action_id"].startswith("set_persona_"): + # Persona selection + return handle_persona_selection(req, client) def view_routing(req: SocketModeRequest, client: SocketModeClient) -> None: diff --git a/backend/danswer/danswerbot/slack/models.py b/backend/danswer/danswerbot/slack/models.py index 57a92a29753..375b92d364c 100644 --- a/backend/danswer/danswerbot/slack/models.py +++ b/backend/danswer/danswerbot/slack/models.py @@ -11,3 +11,4 @@ class SlackMessageInfo(BaseModel): bypass_filters: bool # User has tagged @DanswerBot is_bot_msg: bool # User is using /DanswerBot is_bot_dm: bool # User is direct messaging to DanswerBot + command: str | None # Slash command used by user diff --git a/backend/danswer/danswerbot/slack/utils.py b/backend/danswer/danswerbot/slack/utils.py index 3132aa6f24a..3b6b7a5fd3c 100644 --- a/backend/danswer/danswerbot/slack/utils.py +++ b/backend/danswer/danswerbot/slack/utils.py @@ -172,8 +172,8 @@ def respond_in_thread( ) if response.get("ok"): success = True - except: - pass + except SlackApiError as e: + logger.exception(f"Failed to post message: {e}") if not success: raise RuntimeError(f"Failed to post message: {response}") @@ -275,6 +275,7 @@ def remove_slack_text_interactions(slack_str: str) -> str: slack_str = SlackTextCleaner.replace_links(slack_str) slack_str = SlackTextCleaner.replace_special_catchall(slack_str) slack_str = SlackTextCleaner.add_zero_width_whitespace_after_tag(slack_str) + slack_str = SlackTextCleaner.handle_bold_syntax_for_slack(slack_str) return slack_str diff --git a/backend/danswer/db/engine.py b/backend/danswer/db/engine.py index 4171cf94487..14174f20e6d 100644 --- a/backend/danswer/db/engine.py +++ b/backend/danswer/db/engine.py @@ -59,7 +59,21 @@ def get_sqlalchemy_engine() -> Engine: global _SYNC_ENGINE if _SYNC_ENGINE is None: connection_string = build_connection_string(db_api=SYNC_DB_API) - _SYNC_ENGINE = create_engine(connection_string, pool_size=40, max_overflow=10) + + keepalive_kwargs = { + "keepalives": 1, # Enable TCP Keepalives + "keepalives_idle": 30, # Idle time before keepalive probes + "keepalives_interval": 5, # Interval between keepalive probes + "keepalives_count": 5, # Number of keepalive probes before connection drop + } + + _SYNC_ENGINE = create_engine( + connection_string, + pool_size=40, + max_overflow=10, + pool_pre_ping=True, + connect_args=keepalive_kwargs, + ) return _SYNC_ENGINE diff --git a/backend/danswer/db/models.py b/backend/danswer/db/models.py index 909236a978b..cf1fdaa7ed9 100644 --- a/backend/danswer/db/models.py +++ b/backend/danswer/db/models.py @@ -1108,6 +1108,16 @@ class SlackBotConfig(Base): persona: Mapped[Persona | None] = relationship("Persona") +class UserSlackPersona(Base): + __tablename__ = "user_slack_persona" + + sender_id: Mapped[str] = mapped_column(primary_key=True) + persona_id: Mapped[int | None] = mapped_column( + ForeignKey("persona.id"), nullable=True + ) + persona: Mapped[Persona | None] = relationship("Persona", foreign_keys=[persona_id]) + + class TaskQueueState(Base): # Currently refers to Celery Tasks __tablename__ = "task_queue_jobs" diff --git a/backend/danswer/db/persona.py b/backend/danswer/db/persona.py index 4726cf42637..26292fc9264 100644 --- a/backend/danswer/db/persona.py +++ b/backend/danswer/db/persona.py @@ -9,6 +9,7 @@ from sqlalchemy import or_ from sqlalchemy import select from sqlalchemy import update +from sqlalchemy.orm import joinedload from sqlalchemy.orm import Session from danswer.auth.schemas import UserRole @@ -621,3 +622,14 @@ def delete_persona_by_name( db_session.execute(stmt) db_session.commit() + + +def get_persona_with_docset_and_prompts( + persona_id: int, db_session: Session +) -> Persona | None: + persona = db_session.scalar( + select(Persona) + .options(joinedload(Persona.document_sets), joinedload(Persona.prompts)) + .filter_by(id=persona_id, deleted=False) + ) + return persona diff --git a/backend/danswer/db/users.py b/backend/danswer/db/users.py index f8a3938027f..34e3ada7d4a 100644 --- a/backend/danswer/db/users.py +++ b/backend/danswer/db/users.py @@ -1,9 +1,12 @@ from collections.abc import Sequence +from sqlalchemy import select from sqlalchemy.orm import Session from sqlalchemy.schema import Column +from danswer.db.models import Persona from danswer.db.models import User +from danswer.db.models import UserSlackPersona def list_users(db_session: Session, q: str = "") -> Sequence[User]: @@ -19,3 +22,30 @@ def get_user_by_email(email: str, db_session: Session) -> User | None: user = db_session.query(User).filter(User.email == email).first() # type: ignore return user + + +def fetch_user_slack_persona( + db_session: Session, sender_id: str +) -> UserSlackPersona | None: + return db_session.scalar( + select(UserSlackPersona).where(UserSlackPersona.sender_id == sender_id) + ) + + +def add_user_slack_persona( + db_session: Session, sender_id: str, persona: Persona +) -> None: + user_persona = UserSlackPersona( + sender_id=sender_id, persona_id=persona.id, persona=persona + ) + db_session.add(user_persona) + db_session.commit() + + +def add_slack_persona_for_user( + db_session: Session, persona: Persona, user_slack_persona: UserSlackPersona +) -> None: + user_slack_persona.persona_id = persona.id + user_slack_persona.persona = persona + + db_session.commit() diff --git a/backend/danswer/document_index/vespa/index.py b/backend/danswer/document_index/vespa/index.py index a7892e40f45..75d06da2c10 100644 --- a/backend/danswer/document_index/vespa/index.py +++ b/backend/danswer/document_index/vespa/index.py @@ -18,14 +18,14 @@ import requests from retry import retry +from danswer.configs.app_configs import ENVIRONMENT from danswer.configs.app_configs import LOG_VESPA_TIMING_INFORMATION from danswer.configs.app_configs import VESPA_CONFIG_SERVER_HOST +from danswer.configs.app_configs import VESPA_FEED_HOST +from danswer.configs.app_configs import VESPA_FEED_PORT from danswer.configs.app_configs import VESPA_HOST from danswer.configs.app_configs import VESPA_PORT from danswer.configs.app_configs import VESPA_TENANT_PORT -from danswer.configs.app_configs import VESPA_FEED_HOST -from danswer.configs.app_configs import VESPA_FEED_PORT -from danswer.configs.app_configs import ENVIRONMENT from danswer.configs.chat_configs import DOC_TIME_DECAY from danswer.configs.chat_configs import EDIT_KEYWORD_QUERY from danswer.configs.chat_configs import HYBRID_ALPHA @@ -608,11 +608,10 @@ def _vespa_hit_to_inference_chunk(hit: dict[str, Any]) -> InferenceChunk: for k, v in cast(dict[str, str], source_links_dict_unprocessed).items() } - if 'web' in fields[SOURCE_TYPE]: - score = 1.0 + if "web" in fields[SOURCE_TYPE]: + pass else: - score = hit.get("relevance", 0) - + hit.get("relevance", 0) inference_chunk = InferenceChunk( chunk_id=fields[CHUNK_ID], @@ -624,10 +623,10 @@ def _vespa_hit_to_inference_chunk(hit: dict[str, Any]) -> InferenceChunk: source_type=fields[SOURCE_TYPE], semantic_identifier=fields[SEMANTIC_IDENTIFIER], boost=fields.get(BOOST, 1), - #boost = boost, + # boost = boost, recency_bias=fields.get("matchfeatures", {}).get(RECENCY_BIAS, 1.0), score=hit.get("relevance", 0), - #score = score, + # score = score, hidden=fields.get(HIDDEN, False), primary_owners=fields.get(PRIMARY_OWNERS), secondary_owners=fields.get(SECONDARY_OWNERS), @@ -635,14 +634,14 @@ def _vespa_hit_to_inference_chunk(hit: dict[str, Any]) -> InferenceChunk: match_highlights=match_highlights, updated_at=updated_at, ) - attrs = vars(inference_chunk) - #logger.info(', '.join("%s: %s" % item for item in attrs.items())) + vars(inference_chunk) + # logger.info(', '.join("%s: %s" % item for item in attrs.items())) return inference_chunk def query_vespa_helper(params): - #logger.info("Vespa Query ---> {0}".format(params)) + # logger.info("Vespa Query ---> {0}".format(params)) response = requests.post( SEARCH_ENDPOINT, @@ -690,7 +689,6 @@ def _query_vespa(query_params: Mapping[str, str | int | float]) -> list[Inferenc if "query" in query_params and not cast(str, query_params["query"]).strip(): raise ValueError("No/empty query received") - params = dict( **query_params, **{ @@ -699,21 +697,26 @@ def _query_vespa(query_params: Mapping[str, str | int | float]) -> list[Inferenc if LOG_VESPA_TIMING_INFORMATION else {}, ) - - #All records including web + + # All records including web params["hits"] = 50 filtered_hits_all = query_vespa_helper(params) - #Only Web Records + # Only Web Records and Salesforce KB articles params["hits"] = 10 - params["yql"] = params["yql"] + ' and source_type contains "web"' - filtered_hits_web = query_vespa_helper(params) + params["yql"] = ( + params["yql"] + + ' and (source_type contains "web" or source_type contains "sfkabarticles")' + ) + filtered_hits_web_sf = query_vespa_helper(params) - filtered_hits_final = filtered_hits_web + filtered_hits_all + filtered_hits_final = filtered_hits_web_sf + filtered_hits_all - inference_chunks = [_vespa_hit_to_inference_chunk(hit) for hit in filtered_hits_final] - #inplace sorting based on score - #inference_chunks.sort(key=lambda x: x.score, reverse=True) + inference_chunks = [ + _vespa_hit_to_inference_chunk(hit) for hit in filtered_hits_final + ] + # inplace sorting based on score + # inference_chunks.sort(key=lambda x: x.score, reverse=True) unique_chunks: dict[tuple[str, int], InferenceChunk] = {} for chunk in inference_chunks: diff --git a/backend/danswer/llm/answering/answer.py b/backend/danswer/llm/answering/answer.py index 6a250d02d2a..1aceaa1f8f7 100644 --- a/backend/danswer/llm/answering/answer.py +++ b/backend/danswer/llm/answering/answer.py @@ -1,3 +1,4 @@ +import re from collections.abc import Iterator from typing import cast from uuid import uuid4 @@ -382,6 +383,53 @@ def _raw_output_for_non_explicit_tool_calling_llms( prompt = prompt_builder.build() yield from message_generator_to_string_generator(self.llm.stream(prompt=prompt)) + def _fix_document_references(self, answer: str) -> str: + """ + Searches the input string for DOCUMENT references in any of these forms: + - DOCUMENT (link) + - [DOCUMENT ] (link) + - DOCUMENT + - [DOCUMENT ] + + + and converts them to the proper citation format: + - If a link is provided, returns a linked citation: [[number]](link) + - Otherwise, returns a non-linked citation: [number] + + + However, if an adjacent citation (linked or non-linked) for the same number already follows + immediately (ignoring whitespace), the DOCUMENT reference is not converted (i.e. it is removed) + to avoid duplicate citations. + """ + pattern = r"\[?DOCUMENT\s+(\d+)\]?(?:\s*\((.*?)\))?" + + def replacer(match: re.Match) -> str: + try: + num = int(match.group(1)) + except Exception: + return match.group(0) + + if match.group(2) and match.group(2).strip(): + citation = f"[[{num}]]({match.group(2).strip()})" + else: + citation = f"[{num}]" + + post_text = answer[match.end() :] + adj_pattern = ( + r"^\s*(\[\[\s*" + + re.escape(str(num)) + + r"\s*\]\]\([^)]+\)|\[\s*" + + re.escape(str(num)) + + r"\s*\])" + ) + if re.match(adj_pattern, post_text): + # If an adjacent citation for the same number exists, return an empty string (skip replacement). + return "" + else: + return citation + + return re.sub(pattern, replacer, answer) + @property def processed_streamed_output(self) -> AnswerStream: if self._processed_stream is not None: @@ -465,7 +513,7 @@ def llm_answer(self) -> str: if isinstance(packet, DanswerAnswerPiece) and packet.answer_piece: answer += packet.answer_piece - return answer + return self._fix_document_references(answer) @property def citations(self) -> list[CitationInfo]: diff --git a/backend/danswer/llm/custom_llm.py b/backend/danswer/llm/custom_llm.py index 0b9924b2bed..900dcc693f1 100644 --- a/backend/danswer/llm/custom_llm.py +++ b/backend/danswer/llm/custom_llm.py @@ -1,26 +1,23 @@ import json +import re from collections.abc import Iterator -import re import requests from langchain.schema.language_model import LanguageModelInput from langchain_core.messages import AIMessage from langchain_core.messages import BaseMessage from requests import Timeout -from danswer.configs.model_configs import GEN_AI_API_ENDPOINT -from danswer.configs.model_configs import GEN_AI_IDENTITY_ENDPOINT +from danswer.configs.model_configs import GEN_AI_API_VERSION from danswer.configs.model_configs import GEN_AI_CLIENT_ID from danswer.configs.model_configs import GEN_AI_CLIENT_SECRET -from danswer.configs.model_configs import GEN_AI_API_VERSION +from danswer.configs.model_configs import GEN_AI_IDENTITY_ENDPOINT from danswer.configs.model_configs import GEN_AI_MAX_OUTPUT_TOKENS from danswer.llm.interfaces import LLM +from danswer.llm.interfaces import LLMConfig from danswer.llm.interfaces import ToolChoiceOptions -from danswer.llm.utils import convert_lm_input_to_basic_string from danswer.llm.utils import convert_lm_input_to_prompt from danswer.utils.logger import setup_logger -from danswer.llm.interfaces import LLMConfig -from langchain.prompts.chat import ChatPromptValue logger = setup_logger() @@ -41,51 +38,47 @@ def requires_api_key(self) -> bool: return False def _get_token(self) -> str: - headers = { - 'Content-Type': 'application/x-www-form-urlencoded' - } + headers = {"Content-Type": "application/x-www-form-urlencoded"} data = { - 'client_id': self._client_id, - 'client_secret': self._client_secret, - 'grant_type': 'client_credentials' + "client_id": self._client_id, + "client_secret": self._client_secret, + "grant_type": "client_credentials", } response = requests.post(self._identity_url, headers=headers, data=data) if response.status_code == 200: response_json = response.json() - access_token = response_json.get('access_token') + access_token = response_json.get("access_token") if access_token: return access_token else: - raise ValueError( - "Failed to get access token from the model server" - ) + raise ValueError("Failed to get access token from the model server") else: - print(f"Access token request failed with status code: {response.status_code}") - raise ValueError( - "Failed to get access token from the model server" + print( + f"Access token request failed with status code: {response.status_code}" ) - + raise ValueError("Failed to get access token from the model server") + def __init__( self, # Not used here but you probably want a model server that isn't completely open api_key: str | None, timeout: int, - endpoint: str | None = 'https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-35-turbo/chat/completions?api-version=2023-03-15-preview', + endpoint: str + | None = "https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-4o-mini-2024-07-18/chat/completions?api-version=2024-06-01", identity_url: str | None = GEN_AI_IDENTITY_ENDPOINT, client_id: str | None = GEN_AI_CLIENT_ID, client_secret: str | None = GEN_AI_CLIENT_SECRET, - max_output_tokens: int = GEN_AI_MAX_OUTPUT_TOKENS, + max_output_tokens: int = int(GEN_AI_MAX_OUTPUT_TOKENS), api_version: str | None = GEN_AI_API_VERSION, ): - if not endpoint: raise ValueError( "Cannot point Danswer to a custom LLM server without providing the " "endpoint for the model server." ) - + if not identity_url: raise ValueError( "Cannot point Danswer to a custom LLM server without providing the " @@ -104,7 +97,7 @@ def __init__( "client_secret for the model server." ) - #TODO: implement api versions for endpoints and add those to model + # TODO: implement api versions for endpoints and add those to model # if not api_version: # raise ValueError( # "Cannot point Danswer to a custom LLM server without providing the " @@ -118,9 +111,9 @@ def __init__( self._max_output_tokens = max_output_tokens self._timeout = timeout self.token = self._get_token() - #TODO: Remove hard-coding - self._model_provider = 'custom' - self._model_version = 'gpt-4' + # TODO: Remove hard-coding + self._model_provider = "custom" + self._model_version = "gpt-4" self._temperature = 0.0 self._api_key = api_key @@ -128,7 +121,6 @@ def __init__( self._max_output_tokens = 7000 def _execute(self, input: LanguageModelInput) -> AIMessage: - headers = { "Content-Type": "application/json", "X-UiPath-LlmGateway-RequestedFeature": "ChatWithAssistant", @@ -137,32 +129,35 @@ def _execute(self, input: LanguageModelInput) -> AIMessage: "Authorization": "Bearer " + self.token, } - #print(f"Input: {input}") + # print(f"Input: {input}") chatPrompt = convert_lm_input_to_prompt(input) json_array = [] messages = chatPrompt.to_messages() for msg in messages: mapped_type = self._map_type(msg.type) - json_obj = {"role": mapped_type, "content": self._clean_json_string(msg.content)} + json_obj = { + "role": mapped_type, + "content": self._clean_json_string(msg.content), + } json_array.append(json_obj) - - data = { - "max_tokens": self._max_output_tokens, - "messages": json_array - } + + data = {"max_tokens": self._max_output_tokens, "messages": json_array} try: print(data) with open("requestdata.json", "w") as fp: json.dump(data, fp) - #json_str = json.dumps(data, ensure_ascii=False, indent=4) - #print(f"Request Data: {json_str}") - #json_data = json.loads(json_str) + # json_str = json.dumps(data, ensure_ascii=False, indent=4) + # print(f"Request Data: {json_str}") + # json_data = json.loads(json_str) response = requests.post( - #self._endpoint, headers=headers, data=json_str, timeout=self._timeout - self._endpoint, headers=headers, json=data, timeout=self._timeout + # self._endpoint, headers=headers, data=json_str, timeout=self._timeout + self._endpoint, + headers=headers, + json=data, + timeout=self._timeout, ) except Timeout as error: raise Timeout(f"Model inference to {self._endpoint} timed out") from error @@ -176,23 +171,23 @@ def _execute(self, input: LanguageModelInput) -> AIMessage: raise e message_content = "No response from LLM server" - if data['choices']: - message_content = data['choices'][0]['message']['content'] + if data["choices"]: + message_content = data["choices"][0]["message"]["content"] # print(message_content) return AIMessage(content=message_content) def _clean_json_string(self, input_string): - input_string = re.sub(r'[\\]*"','"', input_string) + input_string = re.sub(r'[\\]*"', '"', input_string) input_string = input_string.replace('"', "'") - + # Remove control characters (ASCII 0-31) - input_string = re.sub(r'[^\x00-\x7F]+', '', input_string) - input_string = re.sub(r'[\xa0]', '', input_string) - + input_string = re.sub(r"[^\x00-\x7F]+", "", input_string) + input_string = re.sub(r"[\xa0]", "", input_string) + # Escape backslashes input_string = input_string.replace("\\", "\\\\") - + return input_string # Convert from AI to LLMGateway types, Only basic, no chunks and no tool and function calls diff --git a/backend/danswer/one_shot_answer/answer_question.py b/backend/danswer/one_shot_answer/answer_question.py index b315c3662c7..3131406cab5 100644 --- a/backend/danswer/one_shot_answer/answer_question.py +++ b/backend/danswer/one_shot_answer/answer_question.py @@ -310,48 +310,67 @@ def get_search_answer( rerank_metrics_callback: Callable[[RerankMetricsContainer], None] | None = None, ) -> OneShotQAResponse: """Collects the streamed one shot answer responses into a single object""" - qa_response = OneShotQAResponse() + max_attempts = 5 + attempt = 0 + qa_response = None - results = stream_answer_objects( - query_req=query_req, - user=user, - max_document_tokens=max_document_tokens, - max_history_tokens=max_history_tokens, - db_session=db_session, - bypass_acl=bypass_acl, - use_citations=use_citations, - danswerbot_flow=danswerbot_flow, - timeout=answer_generation_timeout, - retrieval_metrics_callback=retrieval_metrics_callback, - rerank_metrics_callback=rerank_metrics_callback, - ) + while attempt < max_attempts: + qa_response = OneShotQAResponse() + + results = stream_answer_objects( + query_req=query_req, + user=user, + max_document_tokens=max_document_tokens, + max_history_tokens=max_history_tokens, + db_session=db_session, + bypass_acl=bypass_acl, + use_citations=use_citations, + danswerbot_flow=danswerbot_flow, + timeout=answer_generation_timeout, + retrieval_metrics_callback=retrieval_metrics_callback, + rerank_metrics_callback=rerank_metrics_callback, + ) + + answer = "" + for packet in results: + if isinstance(packet, QueryRephrase): + qa_response.rephrase = packet.rephrased_query + if isinstance(packet, DanswerAnswerPiece) and packet.answer_piece: + answer += packet.answer_piece + elif isinstance(packet, QADocsResponse): + qa_response.docs = packet + elif isinstance(packet, LLMRelevanceFilterResponse): + qa_response.llm_chunks_indices = packet.relevant_chunk_indices + elif isinstance(packet, DanswerQuotes): + qa_response.quotes = packet + elif isinstance(packet, CitationInfo): + if qa_response.citations: + qa_response.citations.append(packet) + else: + qa_response.citations = [packet] + elif isinstance(packet, DanswerContexts): + qa_response.contexts = packet + elif isinstance(packet, StreamingError): + qa_response.error_msg = packet.error + elif isinstance(packet, ChatMessageDetail): + qa_response.chat_message_id = packet.message_id + + if answer: + qa_response.answer = answer + + if use_citations and qa_response.answer and qa_response.citations: + qa_response.answer, qa_response.citations = reorganize_citations( + qa_response.answer, qa_response.citations + ) + break # Citations found, break out of retry loop. + # If citations are not required, we can exit immediately. + elif not use_citations: + break - answer = "" - for packet in results: - if isinstance(packet, QueryRephrase): - qa_response.rephrase = packet.rephrased_query - if isinstance(packet, DanswerAnswerPiece) and packet.answer_piece: - answer += packet.answer_piece - elif isinstance(packet, QADocsResponse): - qa_response.docs = packet - elif isinstance(packet, LLMRelevanceFilterResponse): - qa_response.llm_chunks_indices = packet.relevant_chunk_indices - elif isinstance(packet, DanswerQuotes): - qa_response.quotes = packet - elif isinstance(packet, CitationInfo): - if qa_response.citations: - qa_response.citations.append(packet) - else: - qa_response.citations = [packet] - elif isinstance(packet, DanswerContexts): - qa_response.contexts = packet - elif isinstance(packet, StreamingError): - qa_response.error_msg = packet.error - elif isinstance(packet, ChatMessageDetail): - qa_response.chat_message_id = packet.message_id - - if answer: - qa_response.answer = answer + logger.info( + f"Citations not found, retrying... (attempt {attempt + 1}/{max_attempts})" + ) + attempt += 1 if enable_reflexion: # Because follow up messages are explicitly tagged, we don't need to verify the answer @@ -361,10 +380,4 @@ def get_search_answer( else: qa_response.answer_valid = True - if use_citations and qa_response.answer and qa_response.citations: - # Reorganize citation nums to be in the same order as the answer - qa_response.answer, qa_response.citations = reorganize_citations( - qa_response.answer, qa_response.citations - ) - return qa_response diff --git a/backend/danswer/prompts/chat_prompts.py b/backend/danswer/prompts/chat_prompts.py index bf2a9d8cdc7..2395674290d 100644 --- a/backend/danswer/prompts/chat_prompts.py +++ b/backend/danswer/prompts/chat_prompts.py @@ -2,7 +2,7 @@ from danswer.prompts.constants import QUESTION_PAT REQUIRE_CITATION_STATEMENT = """ -Cite relevant statements INLINE using the format [1], [2], [3], etc to reference the document number, \ +CRUCIAL: Cite relevant statements INLINE using the format [1], [2], [3], etc to reference the document number, \ DO NOT provide a reference section at the end and DO NOT provide any links following the citations. """.rstrip() @@ -11,7 +11,7 @@ """.rstrip() CITATION_REMINDER = """ -Remember to provide inline citations in the format [1], [2], [3], etc. +MANDATORY: For any information sourced from documents in your response, include inline citations formatted as [1], [2], [3], etc. """ ADDITIONAL_INFO = "\n\nAdditional Information:\n\t- {datetime_info}." diff --git a/backend/danswer/prompts/prompt_utils.py b/backend/danswer/prompts/prompt_utils.py index 6d7bddeec95..117de69cf9a 100644 --- a/backend/danswer/prompts/prompt_utils.py +++ b/backend/danswer/prompts/prompt_utils.py @@ -178,9 +178,10 @@ def drop_messages_history_overflow( final_msgs = [final_msg] # Start dropping from the history if necessary - ind_prev_msg_start = find_last_index( - token_counts, max_prompt_tokens=max_allowed_tokens - ) + # ind_prev_msg_start = find_last_index( + # token_counts, max_prompt_tokens=max_allowed_tokens + # ) + ind_prev_msg_start = 0 if system_msg and ind_prev_msg_start <= len(history_msgs): final_messages.append(system_msg) diff --git a/backend/danswer/server/danswer_api/ingestion.py b/backend/danswer/server/danswer_api/ingestion.py index 1b6e6d9852f..66ea99e9f32 100644 --- a/backend/danswer/server/danswer_api/ingestion.py +++ b/backend/danswer/server/danswer_api/ingestion.py @@ -3,6 +3,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.configs.constants import DocumentSource from danswer.connectors.models import Document from danswer.connectors.models import IndexAttemptMetadata @@ -25,7 +26,7 @@ logger = setup_logger() # not using /api to avoid confusion with nginx api path routing -router = APIRouter(prefix="/danswer-api") +router = APIRouter(prefix="/danswer-api", dependencies=[Depends(validate_api_key)]) @router.get("/connector-docs/{cc_pair_id}") diff --git a/backend/danswer/server/documents/cc_pair.py b/backend/danswer/server/documents/cc_pair.py index c6026401d6c..c0c17892239 100644 --- a/backend/danswer/server/documents/cc_pair.py +++ b/backend/danswer/server/documents/cc_pair.py @@ -4,6 +4,7 @@ from sqlalchemy.exc import IntegrityError from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.background.celery.celery_utils import get_deletion_status @@ -19,7 +20,7 @@ from danswer.server.documents.models import ConnectorCredentialPairMetadata from danswer.server.models import StatusResponse -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) @router.get("/admin/cc-pair/{cc_pair_id}") diff --git a/backend/danswer/server/documents/connector.py b/backend/danswer/server/documents/connector.py index ad25523817d..6ca9d23ab90 100644 --- a/backend/danswer/server/documents/connector.py +++ b/backend/danswer/server/documents/connector.py @@ -11,6 +11,7 @@ from pydantic import BaseModel from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.background.celery.celery_utils import get_deletion_status @@ -90,7 +91,7 @@ _GOOGLE_DRIVE_CREDENTIAL_ID_COOKIE_NAME = "google_drive_credential_id" -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) """Admin only API endpoints""" diff --git a/backend/danswer/server/documents/credential.py b/backend/danswer/server/documents/credential.py index a5e9098046a..82641355381 100644 --- a/backend/danswer/server/documents/credential.py +++ b/backend/danswer/server/documents/credential.py @@ -3,6 +3,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.schemas import UserRole from danswer.auth.users import current_admin_user from danswer.auth.users import current_user @@ -19,7 +20,7 @@ from danswer.server.models import StatusResponse -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) """Admin-only endpoints""" diff --git a/backend/danswer/server/documents/document.py b/backend/danswer/server/documents/document.py index 3b0adea246c..059b70758ce 100644 --- a/backend/danswer/server/documents/document.py +++ b/backend/danswer/server/documents/document.py @@ -4,6 +4,7 @@ from fastapi import Query from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_user from danswer.db.embedding_model import get_current_db_embedding_model from danswer.db.engine import get_session @@ -16,7 +17,7 @@ from danswer.server.documents.models import DocumentInfo -router = APIRouter(prefix="/document") +router = APIRouter(prefix="/document", dependencies=[Depends(validate_api_key)]) # Have to use a query parameter as FastAPI is interpreting the URL type document_ids diff --git a/backend/danswer/server/features/document_set/api.py b/backend/danswer/server/features/document_set/api.py index f939329bf9a..3cdaf7b9c21 100644 --- a/backend/danswer/server/features/document_set/api.py +++ b/backend/danswer/server/features/document_set/api.py @@ -3,6 +3,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.db.document_set import check_document_sets_are_public @@ -23,7 +24,7 @@ from danswer.server.features.document_set.models import DocumentSetUpdateRequest -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) @router.post("/admin/document-set") diff --git a/backend/danswer/server/features/folder/api.py b/backend/danswer/server/features/folder/api.py index 000207370d6..754e3693dab 100644 --- a/backend/danswer/server/features/folder/api.py +++ b/backend/danswer/server/features/folder/api.py @@ -4,6 +4,7 @@ from fastapi import Path from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_user from danswer.db.chat import get_chat_session_by_id from danswer.db.engine import get_session @@ -24,7 +25,7 @@ from danswer.server.models import DisplayPriorityRequest from danswer.server.query_and_chat.models import ChatSessionDetails -router = APIRouter(prefix="/folder") +router = APIRouter(prefix="/folder", dependencies=[Depends(validate_api_key)]) @router.get("") diff --git a/backend/danswer/server/features/persona/api.py b/backend/danswer/server/features/persona/api.py index 6739da46606..cd7deb5321a 100644 --- a/backend/danswer/server/features/persona/api.py +++ b/backend/danswer/server/features/persona/api.py @@ -5,6 +5,7 @@ from pydantic import BaseModel from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.db.engine import get_session @@ -27,8 +28,10 @@ logger = setup_logger() -admin_router = APIRouter(prefix="/admin/persona") -basic_router = APIRouter(prefix="/persona") +admin_router = APIRouter( + prefix="/admin/persona", dependencies=[Depends(validate_api_key)] +) +basic_router = APIRouter(prefix="/persona", dependencies=[Depends(validate_api_key)]) class IsVisibleRequest(BaseModel): diff --git a/backend/danswer/server/features/prompt/api.py b/backend/danswer/server/features/prompt/api.py index aebcbb8434d..6a4ddbec18d 100644 --- a/backend/danswer/server/features/prompt/api.py +++ b/backend/danswer/server/features/prompt/api.py @@ -4,6 +4,7 @@ from sqlalchemy.orm import Session from starlette import status +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_user from danswer.db.engine import get_session from danswer.db.models import User @@ -22,7 +23,7 @@ logger = setup_logger() -basic_router = APIRouter(prefix="/prompt") +basic_router = APIRouter(prefix="/prompt", dependencies=[Depends(validate_api_key)]) def create_update_prompt( diff --git a/backend/danswer/server/features/tool/api.py b/backend/danswer/server/features/tool/api.py index b1f57a1a924..f403633ce11 100644 --- a/backend/danswer/server/features/tool/api.py +++ b/backend/danswer/server/features/tool/api.py @@ -6,6 +6,7 @@ from pydantic import BaseModel from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.db.engine import get_session @@ -20,8 +21,8 @@ from danswer.tools.custom.openapi_parsing import openapi_to_method_specs from danswer.tools.custom.openapi_parsing import validate_openapi_schema -router = APIRouter(prefix="/tool") -admin_router = APIRouter(prefix="/admin/tool") +router = APIRouter(prefix="/tool", dependencies=[Depends(validate_api_key)]) +admin_router = APIRouter(prefix="/admin/tool", dependencies=[Depends(validate_api_key)]) class CustomToolCreate(BaseModel): diff --git a/backend/danswer/server/gpts/api.py b/backend/danswer/server/gpts/api.py index 84b0078ee77..3be9c47891e 100644 --- a/backend/danswer/server/gpts/api.py +++ b/backend/danswer/server/gpts/api.py @@ -6,6 +6,7 @@ from pydantic import BaseModel from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.db.engine import get_session from danswer.llm.factory import get_default_llms from danswer.search.models import SearchRequest @@ -13,11 +14,10 @@ from danswer.server.danswer_api.ingestion import api_key_dep from danswer.utils.logger import setup_logger - logger = setup_logger() -router = APIRouter(prefix="/gpts") +router = APIRouter(prefix="/gpts", dependencies=[Depends(validate_api_key)]) def time_ago(dt: datetime) -> str: diff --git a/backend/danswer/server/manage/administrative.py b/backend/danswer/server/manage/administrative.py index d6a52917f3b..84d0b390f57 100644 --- a/backend/danswer/server/manage/administrative.py +++ b/backend/danswer/server/manage/administrative.py @@ -8,6 +8,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.configs.app_configs import GENERATIVE_MODEL_ACCESS_CHECK_FREQ from danswer.configs.constants import DocumentSource @@ -32,7 +33,7 @@ from danswer.server.manage.models import HiddenUpdateRequest from danswer.utils.logger import setup_logger -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) logger = setup_logger() GEN_AI_KEY_CHECK_TIME = "genai_api_key_last_check_time" diff --git a/backend/danswer/server/manage/get_state.py b/backend/danswer/server/manage/get_state.py index 3ca47841b64..5d61e4f4376 100644 --- a/backend/danswer/server/manage/get_state.py +++ b/backend/danswer/server/manage/get_state.py @@ -1,13 +1,15 @@ from fastapi import APIRouter +from fastapi import Depends from danswer import __version__ +from danswer.auth.api_key import validate_api_key from danswer.auth.users import user_needs_to_be_verified from danswer.configs.app_configs import AUTH_TYPE from danswer.server.manage.models import AuthTypeResponse from danswer.server.manage.models import VersionResponse from danswer.server.models import StatusResponse -router = APIRouter() +router = APIRouter(dependencies=[Depends(validate_api_key)]) @router.get("/health") diff --git a/backend/danswer/server/manage/llm/api.py b/backend/danswer/server/manage/llm/api.py index 4df00b529af..71047ac36e2 100644 --- a/backend/danswer/server/manage/llm/api.py +++ b/backend/danswer/server/manage/llm/api.py @@ -5,6 +5,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.db.engine import get_session @@ -28,8 +29,8 @@ logger = setup_logger() -admin_router = APIRouter(prefix="/admin/llm") -basic_router = APIRouter(prefix="/llm") +admin_router = APIRouter(prefix="/admin/llm", dependencies=[Depends(validate_api_key)]) +basic_router = APIRouter(prefix="/llm", dependencies=[Depends(validate_api_key)]) @admin_router.get("/built-in/options") diff --git a/backend/danswer/server/manage/secondary_index.py b/backend/danswer/server/manage/secondary_index.py index 6f5adf752f6..f52b944dbab 100644 --- a/backend/danswer/server/manage/secondary_index.py +++ b/backend/danswer/server/manage/secondary_index.py @@ -4,6 +4,7 @@ from fastapi import status from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.configs.app_configs import DISABLE_INDEX_UPDATE_ON_SWAP @@ -23,7 +24,7 @@ from danswer.server.models import IdReturn from danswer.utils.logger import setup_logger -router = APIRouter(prefix="/secondary-index") +router = APIRouter(prefix="/secondary-index", dependencies=[Depends(validate_api_key)]) logger = setup_logger() diff --git a/backend/danswer/server/manage/slack_bot.py b/backend/danswer/server/manage/slack_bot.py index 71ee4df1f27..ea624e47f63 100644 --- a/backend/danswer/server/manage/slack_bot.py +++ b/backend/danswer/server/manage/slack_bot.py @@ -3,6 +3,7 @@ from fastapi import HTTPException from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.danswerbot.slack.config import validate_channel_names from danswer.danswerbot.slack.tokens import fetch_tokens @@ -24,7 +25,7 @@ from danswer.server.manage.models import SlackBotTokens -router = APIRouter(prefix="/manage") +router = APIRouter(prefix="/manage", dependencies=[Depends(validate_api_key)]) def _form_channel_config( @@ -73,7 +74,9 @@ def _form_channel_config( if respond_tag_only is not None: channel_config["respond_tag_only"] = respond_tag_only if respond_team_member_list: - channel_config["respond_team_member_list"] = [item.lower() for item in respond_team_member_list] + channel_config["respond_team_member_list"] = [ + item.lower() for item in respond_team_member_list + ] if respond_slack_group_list: channel_config["respond_slack_group_list"] = respond_slack_group_list if answer_filters: diff --git a/backend/danswer/server/manage/users.py b/backend/danswer/server/manage/users.py index c635469919e..8e505755861 100644 --- a/backend/danswer/server/manage/users.py +++ b/backend/danswer/server/manage/users.py @@ -9,6 +9,7 @@ from sqlalchemy import update from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.invited_users import get_invited_users from danswer.auth.invited_users import write_invited_users from danswer.auth.noauth_user import fetch_no_auth_user @@ -38,7 +39,7 @@ logger = setup_logger() -router = APIRouter() +router = APIRouter(dependencies=[Depends(validate_api_key)]) USERS_PAGE_SIZE = 10 diff --git a/backend/danswer/server/query_and_chat/chat_backend.py b/backend/danswer/server/query_and_chat/chat_backend.py index 646660c9fac..f03dcb93397 100644 --- a/backend/danswer/server/query_and_chat/chat_backend.py +++ b/backend/danswer/server/query_and_chat/chat_backend.py @@ -11,6 +11,7 @@ from pydantic import BaseModel from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_user from danswer.chat.chat_utils import create_chat_chain from danswer.chat.process_message import stream_chat_message @@ -69,7 +70,8 @@ logger = setup_logger() -router = APIRouter(prefix="/chat") +router = APIRouter(prefix="/chat", dependencies=[Depends(validate_api_key)]) +# api_router = APIRouter(prefix="/chat", dependencies=[Depends(validate_api_key)]) @router.get("/get-user-chat-sessions") diff --git a/backend/danswer/server/query_and_chat/query_backend.py b/backend/danswer/server/query_and_chat/query_backend.py index 43192211b79..ff632e0613a 100644 --- a/backend/danswer/server/query_and_chat/query_backend.py +++ b/backend/danswer/server/query_and_chat/query_backend.py @@ -4,6 +4,7 @@ from fastapi.responses import StreamingResponse from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.configs.constants import DocumentSource @@ -34,8 +35,8 @@ logger = setup_logger() -admin_router = APIRouter(prefix="/admin") -basic_router = APIRouter(prefix="/query") +admin_router = APIRouter(prefix="/admin", dependencies=[Depends(validate_api_key)]) +basic_router = APIRouter(prefix="/query", dependencies=[Depends(validate_api_key)]) @admin_router.post("/search") diff --git a/backend/danswer/server/settings/api.py b/backend/danswer/server/settings/api.py index 422e268c13e..953e6e0ae4b 100644 --- a/backend/danswer/server/settings/api.py +++ b/backend/danswer/server/settings/api.py @@ -2,6 +2,7 @@ from fastapi import Depends from fastapi import HTTPException +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.auth.users import current_user from danswer.db.models import User @@ -10,8 +11,10 @@ from danswer.server.settings.store import store_settings -admin_router = APIRouter(prefix="/admin/settings") -basic_router = APIRouter(prefix="/settings") +admin_router = APIRouter( + prefix="/admin/settings", dependencies=[Depends(validate_api_key)] +) +basic_router = APIRouter(prefix="/settings", dependencies=[Depends(validate_api_key)]) @admin_router.put("") diff --git a/backend/danswer/server/token_rate_limits/api.py b/backend/danswer/server/token_rate_limits/api.py index 245e3391410..093a2c84848 100644 --- a/backend/danswer/server/token_rate_limits/api.py +++ b/backend/danswer/server/token_rate_limits/api.py @@ -2,6 +2,7 @@ from fastapi import Depends from sqlalchemy.orm import Session +from danswer.auth.api_key import validate_api_key from danswer.auth.users import current_admin_user from danswer.db.engine import get_session from danswer.db.models import User @@ -13,7 +14,9 @@ from ee.danswer.db.token_limit import insert_global_token_rate_limit from ee.danswer.db.token_limit import update_token_rate_limit -router = APIRouter(prefix="/admin/token-rate-limits") +router = APIRouter( + prefix="/admin/token-rate-limits", dependencies=[Depends(validate_api_key)] +) """ diff --git a/backend/danswer/utils/text_processing.py b/backend/danswer/utils/text_processing.py index b0fbcdfa1e9..67188e863b9 100644 --- a/backend/danswer/utils/text_processing.py +++ b/backend/danswer/utils/text_processing.py @@ -20,7 +20,18 @@ def decode_escapes(s: str) -> str: def decode_match(match: re.Match) -> str: - return codecs.decode(match.group(0), "unicode-escape") + matched_str = match.group(0) + + # Only double escape non-Unicode sequences + if matched_str.startswith("\\U") and not re.match( + r"\\U[0-9a-fA-F]{8}", matched_str + ): + return matched_str + + try: + return codecs.decode(matched_str, "unicode-escape") + except UnicodeDecodeError: + return matched_str return ESCAPE_SEQUENCE_RE.sub(decode_match, s) diff --git a/darwin-kubernetes/env-configmap.yaml b/darwin-kubernetes/env-configmap.yaml index e796c046e94..d5dd301ffed 100644 --- a/darwin-kubernetes/env-configmap.yaml +++ b/darwin-kubernetes/env-configmap.yaml @@ -15,7 +15,7 @@ data: EMAIL_FROM: "" # 'your-email@company.com' SMTP_USER missing used instead # Gen AI Settings GEN_AI_MODEL_PROVIDER: "custom" - GEN_AI_API_ENDPOINT: "https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-35-turbo/chat/completions?api-version=2023-03-15-preview" + GEN_AI_API_ENDPOINT: "https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-4o-mini-2024-07-18/chat/completions?api-version=2024-06-01" GEN_AI_IDENTITY_ENDPOINT: "https://alpha.uipath.com/identity_/connect/token" GEN_AI_CLIENT_ID: "XXX" GEN_AI_CLIENT_SECRET: "XXX" diff --git a/deployment/docker_compose/docker-compose.local.yml b/deployment/docker_compose/docker-compose.local.yml index 7583b09b836..99179c0747f 100644 --- a/deployment/docker_compose/docker-compose.local.yml +++ b/deployment/docker_compose/docker-compose.local.yml @@ -39,7 +39,7 @@ services: - GEN_AI_MODEL_VERSION=${GEN_AI_MODEL_VERSION:-} - FAST_GEN_AI_MODEL_VERSION=${FAST_GEN_AI_MODEL_VERSION:-} - GEN_AI_API_KEY=${GEN_AI_API_KEY:-} - - GEN_AI_API_ENDPOINT=https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-35-turbo/chat/completions?api-version=2023-03-15-preview + - GEN_AI_API_ENDPOINT=https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-4o-mini-2024-07-18/chat/completions?api-version=2024-06-01 - GEN_AI_IDENTITY_ENDPOINT=https://alpha.uipath.com/identity_/connect/token - GEN_AI_CLIENT_ID=${GEN_AI_CLIENT_ID:-} - GEN_AI_CLIENT_SECRET=${GEN_AI_CLIENT_SECRET:-} @@ -128,7 +128,7 @@ services: - GEN_AI_MODEL_VERSION=${GEN_AI_MODEL_VERSION:-} - FAST_GEN_AI_MODEL_VERSION=${FAST_GEN_AI_MODEL_VERSION:-custom} - GEN_AI_API_KEY=${GEN_AI_API_KEY:-} - - GEN_AI_API_ENDPOINT=https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-35-turbo/chat/completions?api-version=2023-03-15-preview + - GEN_AI_API_ENDPOINT=https://alpha.uipath.com/llmgateway_/openai/deployments/gpt-4o-mini-2024-07-18/chat/completions?api-version=2024-06-01 - GEN_AI_IDENTITY_ENDPOINT=https://alpha.uipath.com/identity_/connect/token - GEN_AI_CLIENT_ID=${GEN_AI_CLIENT_ID:-} - GEN_AI_CLIENT_SECRET=${GEN_AI_CLIENT_SECRET:-} diff --git a/web/src/app/admin/connectors/jira/page.tsx b/web/src/app/admin/connectors/jira/page.tsx index f960348e6da..588a8db5ab5 100644 --- a/web/src/app/admin/connectors/jira/page.tsx +++ b/web/src/app/admin/connectors/jira/page.tsx @@ -26,18 +26,6 @@ import { usePublicCredentials } from "@/lib/hooks"; import { AdminPageTitle } from "@/components/admin/Title"; import { Card, Divider, Text, Title } from "@tremor/react"; -// Copied from the `extract_jira_project` function -const extractJiraProject = (url: string): string | null => { - const parsedUrl = new URL(url); - const splitPath = parsedUrl.pathname.split("/"); - const projectPos = splitPath.indexOf("projects"); - if (projectPos !== -1 && splitPath.length > projectPos + 1) { - const jiraProject = splitPath[projectPos + 1]; - return jiraProject; - } - return null; -}; - const Main = () => { const { popup, setPopup } = usePopup(); @@ -224,14 +212,8 @@ const Main = () => { <> {" "} - Specify any link to a Jira page below and click "Index" to - Index. Based on the provided link, we will index the ENTIRE PROJECT, - not just the specified page. For example, entering{" "} - - https://danswer.atlassian.net/jira/software/projects/DAN/boards/1 - {" "} - and clicking the Index button will index the whole DAN Jira - project. + Please specify the filters you want to use for indexing the Jira + Issues. {jiraConnectorIndexingStatuses.length > 0 && ( <> @@ -261,19 +243,12 @@ const Main = () => { }} specialColumns={[ { - header: "Url", - key: "url", + header: "Filters", + key: "filters", getValue: (ccPairStatus) => { const connectorConfig = ccPairStatus.connector.connector_specific_config; - return ( - - {connectorConfig.jira_project_url} - - ); + return connectorConfig.jira_filter; }, }, { @@ -300,21 +275,15 @@ const Main = () => {

Add a New Project

- nameBuilder={(values) => - `JiraConnector-${values.jira_project_url}` - } - ccPairNameBuilder={(values) => - extractJiraProject(values.jira_project_url) - } + nameBuilder={(values) => `JiraConnector-${values.jira_filter}`} + ccPairNameBuilder={(values) => `JIRA - ${values.jira_filter}`} credentialId={jiraCredential.id} source="jira" inputType="poll" formBody={ <> - + + } formBodyBuilder={(values) => { @@ -331,15 +300,19 @@ const Main = () => { ); }} validationSchema={Yup.object().shape({ - jira_project_url: Yup.string().required( - "Please enter any link to your jira project e.g. https://danswer.atlassian.net/jira/software/projects/DAN/boards/1" + jira_base_url: Yup.string().required( + "Please provide the base url e.g. https://danswer.atlassian.net" + ), + jira_filter: Yup.string().required( + "Please provide some filters." ), comment_email_blacklist: Yup.array() .of(Yup.string().required("Emails names must be strings")) .required(), })} initialValues={{ - jira_project_url: "", + jira_base_url: "", + jira_filter: "", comment_email_blacklist: [], }} refreshFreq={10 * 60} // 10 minutes diff --git a/web/src/app/admin/connectors/sfkbarticles/page.tsx b/web/src/app/admin/connectors/sfkbarticles/page.tsx new file mode 100644 index 00000000000..d333638e7a0 --- /dev/null +++ b/web/src/app/admin/connectors/sfkbarticles/page.tsx @@ -0,0 +1,290 @@ +"use client"; + +import * as Yup from "yup"; +import { TrashIcon, SalesforceIcon } from "@/components/icons/icons"; // Make sure you have a Document360 icon +import { errorHandlingFetcher as fetcher } from "@/lib/fetcher"; +import useSWR, { useSWRConfig } from "swr"; +import { LoadingAnimation } from "@/components/Loading"; +import { HealthCheckBanner } from "@/components/health/healthcheck"; +import { + SfKbArticlesConfig, + SfKbArticlesCredentialJson, + ConnectorIndexingStatus, + Credential, +} from "@/lib/types"; // Modify or create these types as required +import { adminDeleteCredential, linkCredential } from "@/lib/credential"; +import { CredentialForm } from "@/components/admin/connectors/CredentialForm"; +import { + TextFormField, + TextArrayFieldBuilder, +} from "@/components/admin/connectors/Field"; +import { ConnectorsTable } from "@/components/admin/connectors/table/ConnectorsTable"; +import { ConnectorForm } from "@/components/admin/connectors/ConnectorForm"; +import { usePublicCredentials } from "@/lib/hooks"; +import { AdminPageTitle } from "@/components/admin/Title"; +import { Card, Text, Title } from "@tremor/react"; + +const MainSection = () => { + const { mutate } = useSWRConfig(); + const { + data: connectorIndexingStatuses, + isLoading: isConnectorIndexingStatusesLoading, + error: isConnectorIndexingStatusesError, + } = useSWR[]>( + "/api/manage/admin/connector/indexing-status", + fetcher + ); + + const { + data: credentialsData, + isLoading: isCredentialsLoading, + error: isCredentialsError, + refreshCredentials, + } = usePublicCredentials(); + + if ( + (!connectorIndexingStatuses && isConnectorIndexingStatusesLoading) || + (!credentialsData && isCredentialsLoading) + ) { + return ; + } + + if (isConnectorIndexingStatusesError || !connectorIndexingStatuses) { + return
Failed to load connectors
; + } + + if (isCredentialsError || !credentialsData) { + return
Failed to load credentials
; + } + + const SalesforceConnectorIndexingStatuses: ConnectorIndexingStatus< + SfKbArticlesConfig, + SfKbArticlesCredentialJson + >[] = connectorIndexingStatuses.filter( + (connectorIndexingStatus) => + connectorIndexingStatus.connector.source === "salesforce" + ); + + const SfKbArticlesCredential: + | Credential + | undefined = credentialsData.find( + (credential) => credential.credential_json?.sf_username + ); + + return ( + <> + + The Salesforce Knowledge Base Articles connector allows you to index and + search through your Salesforce Knowledge Base. Once setup, all indicated + Salesforce data will be queryable within Darwin. + + + + Step 1: Provide Salesforce credentials + + {SfKbArticlesCredential ? ( + <> +
+ Existing SalesForce Username: + + {SfKbArticlesCredential.credential_json.sf_username} + + +
+ + ) : ( + <> + + As a first step, please provide the Salesforce account's + client_id, client_secret, username and password. + + + + formBody={ + <> + + + + + + } + validationSchema={Yup.object().shape({ + sf_client_id: Yup.string().required( + "Please enter your Salesforce Client Id" + ), + sf_client_secret: Yup.string().required( + "Please enter your Salesforce Client Secret" + ), + sf_username: Yup.string().required( + "Please enter your Salesforce username" + ), + sf_password: Yup.string().required( + "Please enter your Salesforce password" + ), + })} + initialValues={{ + sf_client_id: "", + sf_client_secret: "", + sf_username: "", + sf_password: "", + }} + onSubmit={(isSuccess) => { + if (isSuccess) { + refreshCredentials(); + } + }} + /> + + + )} + + + Step 2: Manage Salesforce KB Articles Connector + + + {SalesforceConnectorIndexingStatuses.length > 0 && ( + <> + + The latest state of your Salesforce objects are fetched every 10 + minutes. + +
+ + connectorIndexingStatuses={SalesforceConnectorIndexingStatuses} + liveCredential={SfKbArticlesCredential} + getCredential={(credential) => + credential.credential_json.sf_client_secret + } + onUpdate={() => + mutate("/api/manage/admin/connector/indexing-status") + } + onCredentialLink={async (connectorId) => { + if (SfKbArticlesCredential) { + await linkCredential(connectorId, SfKbArticlesCredential.id); + mutate("/api/manage/admin/connector/indexing-status"); + } + }} + specialColumns={[ + { + header: "Connectors", + key: "connectors", + getValue: (ccPairStatus) => { + const connectorConfig = + ccPairStatus.connector.connector_specific_config; + return `${connectorConfig.requested_objects}`; + }, + }, + ]} + includeName + /> +
+ + )} + + {SfKbArticlesCredential ? ( + + + nameBuilder={(values) => + values.requested_objects && values.requested_objects.length > 0 + ? `SfKbArticles-${values.requested_objects.join("-")}` + : "SfKbArticles" + } + ccPairNameBuilder={(values) => + values.requested_objects && values.requested_objects.length > 0 + ? `SfKbArticles-${values.requested_objects.join("-")}` + : "SfKbArticles" + } + source="sfkbarticles" + inputType="poll" + // formBody={<>} + formBodyBuilder={TextArrayFieldBuilder({ + name: "requested_objects", + label: "Specify the Product Components", + subtext: ( + <> +
+ Specify the product components for which you want to fetch the + Salesforce Knowledge Base articles. +
+
+ Example:{" "} + + Orchestrator, Activities, Studio, Robot, Automation Hub + + . +
+
+ By default, it will fetch articles for all the product + components. +
+
+ Hint: Use the exact product component name for accurate + results. + + ), + })} + validationSchema={Yup.object().shape({ + requested_objects: Yup.array() + .of( + Yup.string().required( + "Salesforce Product Component names must be strings" + ) + ) + .required(), + })} + initialValues={{ + requested_objects: [], + }} + credentialId={SfKbArticlesCredential.id} + refreshFreq={10 * 60} // 10 minutes + /> +
+ ) : ( + + Please provide all Salesforce info in Step 1 first! Once you're + done with that, you can then specify the product components for which + you want to fetch the Salesforce Knowledge Base articles. + + )} + + ); +}; + +export default function Page() { + return ( +
+
+ +
+ + } + title="Salesforce KB Articles" + /> + + +
+ ); +} diff --git a/web/src/components/admin/connectors/ConnectorTitle.tsx b/web/src/components/admin/connectors/ConnectorTitle.tsx index a0065025921..01f357c4118 100644 --- a/web/src/components/admin/connectors/ConnectorTitle.tsx +++ b/web/src/components/admin/connectors/ConnectorTitle.tsx @@ -54,8 +54,8 @@ export const ConnectorTitle = ({ } else if (connector.source === "jira") { const typedConnector = connector as Connector; additionalMetadata.set( - "Jira Project URL", - typedConnector.connector_specific_config.jira_project_url + "Jira Filters", + typedConnector.connector_specific_config.jira_filter ); } else if (connector.source === "google_drive") { const typedConnector = connector as Connector; diff --git a/web/src/lib/sources.ts b/web/src/lib/sources.ts index f141e42edb5..6c93d412a69 100644 --- a/web/src/lib/sources.ts +++ b/web/src/lib/sources.ts @@ -172,6 +172,11 @@ const SOURCE_METADATA_MAP: SourceMap = { displayName: "Salesforce", category: SourceCategory.AppConnection, }, + sfkbarticles: { + icon: SalesforceIcon, + displayName: "SfKbArticles", + category: SourceCategory.AppConnection, + }, sharepoint: { icon: SharepointIcon, displayName: "Sharepoint", diff --git a/web/src/lib/types.ts b/web/src/lib/types.ts index fe252afbf28..0344b002ccf 100644 --- a/web/src/lib/types.ts +++ b/web/src/lib/types.ts @@ -51,6 +51,7 @@ export type ValidSources = | "loopio" | "dropbox" | "salesforce" + | "sfkbarticles" | "sharepoint" | "teams" | "zendesk" @@ -136,7 +137,8 @@ export interface ConfluenceConfig { } export interface JiraConfig { - jira_project_url: string; + jira_base_url: string; + jira_filter: string; comment_email_blacklist?: string[]; } @@ -144,6 +146,10 @@ export interface SalesforceConfig { requested_objects?: string[]; } +export interface SfKbArticlesConfig { + requested_objects?: string[]; +} + export interface SharepointConfig { sites?: string[]; } @@ -451,12 +457,20 @@ export interface OCICredentialJson { access_key_id: string; secret_access_key: string; } + export interface SalesforceCredentialJson { sf_username: string; sf_password: string; sf_security_token: string; } +export interface SfKbArticlesCredentialJson { + sf_client_id: string; + sf_client_secret: string; + sf_username: string; + sf_password: string; +} + export interface SharepointCredentialJson { sp_client_id: string; sp_client_secret: string;