
    j8                         d dl Z d dlmZmZmZmZ d dlmZ d dlm	Z	 d dl
mZ d dlmZmZmZmZ  G d de	      Z G d	 d
e      Z	 ddededee   dedee   defdZ G d d      Z G d d      Zy)    N)ListOptionalLiteralTuple)BaseLLMModelInfo)BaseLLMException)AllMessageValues)ChatCompletionToolCallChunkChatCompletionUsageBlockGenericStreamingChunkProviderSpecificModelInfoc                        e Zd Z fdZ xZS )CohereErrorc                 (    t         |   ||       y )N)status_codemessage)super__init__)selfr   r   	__class__s      y/Users/manta/Documents/Projects/TheRoad-I1/backend/.venv/lib/python3.12/site-packages/litellm/llms/cohere/common_utils.pyr   zCohereError.__init__   s    ['B    )__name__
__module____qualname__r   __classcell__)r   s   @r   r   r      s    C Cr   r   c                      e Zd Zdedee   fdZ	 ddee   dee   dee   fdZe	ddee   dee   fd       Z
e		 ddee   dee   fd	       Z	 	 dd
ededee   dededee   dee   defdZe	dedee   fd       Ze	deded   fd       Zy)CohereModelInfomodelreturnc                      y)zE
        Default values all models of this provider support.
        N )r   r   s     r   get_provider_infoz!CohereModelInfo.get_provider_info   s     r   Napi_keyapi_basec                     g S )zF
        Returns a list of models supported by this provider.
        r"   )r   r$   r%   s      r   
get_modelszCohereModelInfo.get_models   s	     	r   c                     | S Nr"   )r$   s    r   get_api_keyzCohereModelInfo.get_api_key&   s    r   c                     | S r)   r"   )r%   s    r   get_api_basezCohereModelInfo.get_api_base*   s	     r   headersmessagesoptional_paramslitellm_paramsc                     i S r)   r"   )r   r-   r   r.   r/   r0   r$   r%   s           r   validate_environmentz$CohereModelInfo.validate_environment0   s	     	r   c                      y)a2  
        Returns the base model name from the given model name.

        Some providers like bedrock - can receive model=`invoke/anthropic.claude-3-opus-20240229-v1:0` or `converse/anthropic.claude-3-opus-20240229-v1:0`
            This function will return `anthropic.claude-3-opus-20240229-v1:0`
        Nr"   r   s    r   get_base_modelzCohereModelInfo.get_base_model<   s     	r   )v1v2c                     d| v ryy)a  
        Get the Cohere route for the given model.
        
        Args:
            model: The model name (e.g., "cohere_chat/v2/command-r-plus", "command-r-plus")
            
        Returns:
            "v2" for standard Cohere v2 API (default), "v1" for Cohere v1 API
        zv1/r6   r7   r"   r4   s    r   get_cohere_routez CohereModelInfo.get_cohere_routeF   s     E> r   )NNr)   )r   r   r   strr   r   r#   r   r'   staticmethodr*   r,   dictr	   r2   r5   r   r9   r"   r   r   r   r      s^    
+	, HL}7?}	c Xc] hsm   "&3-	#  "&"&

 
 '(	

 
 
 #
 3-
 

 c hsm    
(;  r   r   r-   r   r.   r/   r$   r    c                 D    | j                  dddd       |rd| | d<   | S )aQ  
    Return headers to use for cohere chat completion request

    Cohere API Ref: https://docs.cohere.com/reference/chat
    Expected headers:
    {
        "Request-Source": "unspecified:litellm",
        "accept": "application/json",
        "content-type": "application/json",
        "Authorization": "Bearer $CO_API_KEY"
    }
    zunspecified:litellmzapplication/json)zRequest-Sourceacceptzcontent-typezBearer Authorization)update)r-   r   r.   r/   r$   s        r   r2   r2   X   s9    & NN3(.	
 %,WI#6 Nr   c                   ^    e Zd Z	 ddedee   fdZdedefdZd Z	d Z
dedefd	Zd
 Zd Zy)ModelResponseIteratorsync_stream	json_modec                 ^    || _         | j                   | _        g | _        d| _        || _        y Nstreaming_responseresponse_iteratorcontent_blocks
tool_indexrD   r   rI   rC   rD   s       r   r   zModelResponseIterator.__init__x   0     #5!%!8!8$&"r   chunkr    c           	         	 d}d }d}d}d }d }t        |j                  dd            }d|v r|d   }nd|v r|d   du r
|d   }|d   }d	|v rd	|d	   i}t        |||||||
      }	|	S # t        j                  $ r t        d|       w xY w)N Findexr   textis_finishedTfinish_reason	citationsrS   tool_userT   rU   usagerR   provider_specific_fieldsz"Failed to decode JSON from chunk: )intgetr   jsonJSONDecodeError
ValueError)
r   rO   rS   rX   rT   rU   rY   rZ   rR   returned_chunks
             r   chunk_parserz"ModelResponseIterator.chunk_parser   s     	KD>BHKM8<E'+$		'1-.EV}%'E-,@D,H#M2 %o 6e#,7{9K+L(2!'+)AN "!## 	KA%IJJ	Ks   A$A' '"B	c                     | S r)   r"   r   s    r   __iter__zModelResponseIterator.__iter__       r   c                    	 | j                   j                         }	 | j                  |      S # t        $ r t        t        $ r}t	        d|       d }~ww xY w# t        $ r t        t        $ r}t	        d| d|       d }~ww xY wNz#Error receiving chunk from stream: rO   zError parsing chunk: z,
Received chunk: rJ   __next__StopIterationr_   RuntimeError"convert_str_chunk_to_generic_chunkr   rO   es      r   rj   zModelResponseIterator.__next__       	J**335E	V:::GG  	  	J!DQCHII	J
  	  	V!6qc9LUGTUU	V,   / A AAAB	3BB	c                     |}t        |t              r,|j                  d      }|j                  d      }|dk7  r||d }t	        j
                  |      }| j                  |      S )z
        Convert a string chunk to a GenericStreamingChunk

        Note: This is used for Cohere pass through streaming logging
        utf-8data:rG   Nrh   
isinstancebytesdecodefindr]   loadsra   r   rO   str_linerR   	data_jsons        r   rm   z8ModelResponseIterator.convert_str_chunk_to_generic_chunk   d     eU#||G,HMM'*E{#EF+JJx(	  y 11r   c                 D    | j                   j                         | _        | S r)   rI   	__aiter__async_response_iteratorrc   s    r   r   zModelResponseIterator.__aiter__       '+'>'>'H'H'J$r   c                 4  K   	 | j                   j                          d {   }	 | j                  |      S 7 # t        $ r t        t        $ r}t	        d|       d }~ww xY w# t        $ r t        t        $ r}t	        d| d|       d }~ww xY wwrg   r   	__anext__StopAsyncIterationr_   rl   rm   rn   s      r   r   zModelResponseIterator.__anext__        	J66@@BBE	V:::GG C! 	%$$ 	J!DQCHII	J
 " 	%$$ 	V!6qc9LUGTUU	VN   B; 9; A( B; A%A  A%%B(B?BBBNF)r   r   r   boolr   r   r<   r   ra   rd   rj   r:   rm   r   r   r"   r   r   rB   rB   w   s`    QV#/3#@H#!K$ !K+@ !KHV2 2@U 2"Vr   rB   c                       e Zd ZdZ	 ddedee   fdZdedefdZ	dedee
   fdZdedee   fd	Zdedee   fd
Zdedeeeee   f   fdZdedefdZd Zd ZdedefdZd Zd Zy)CohereV2ModelResponseIteratorz2V2-specific response iterator for Cohere streamingrC   rD   c                 ^    || _         | j                   | _        g | _        d| _        || _        y rF   rH   rM   s       r   r   z&CohereV2ModelResponseIterator.__init__   rN   r   rO   r    c                     |j                  di       }|j                  di       }|j                  di       }t        |t              r	d|v r|d   S t        |t              r|S y)z+Parse content-delta chunks to extract text.deltar   contentrS   rQ   )r\   rv   r<   r:   )r   rO   r   r   r   s        r   _parse_content_deltaz2CohereV2ModelResponseIterator._parse_content_delta   s_    		'2&))Ir*++i,gt$7):6?"%Nr   c                     |j                  di       }|j                  dg       }|rB|d   j                  dd      d|d   j                  dd      |d   j                  dd      d	d
S y)z3Parse tool-call-delta chunks to extract tool calls.r   
tool_callsr   idrQ   functionname	arguments)r   r   )r   typer   Nr\   )r   rO   r   r   s       r   _parse_tool_call_deltaz4CohereV2ModelResponseIterator._parse_tool_call_delta   sx    		'2&YY|R0
 m''b1"&qM--fb9!+A!2!2;!C  r   c                     |j                  di       }|j                  di       }|j                  di       }|j                  dd      }|rd|iS y)z2Parse tool-plan-delta events to extract tool plan.datar   r   	tool_planrQ   Nr   )r   rO   r   r   r   r   s         r   _parse_tool_plan_deltaz4CohereV2ModelResponseIterator._parse_tool_plan_delta  sS    yy$"%))Ir*KKR0	++r   c                 R   |j                  di       }|j                  di       }|j                  di       }|j                  di       }|r]|j                  dd      |j                  dd      |j                  dd	      |j                  d
g       |j                  dd      d}d|giS y)z1Parse citation-start events to extract citations.r   r   r   rV   startr   endrS   rQ   sourcesr   TEXT_CONTENT)r   r   rS   r   r   Nr   )r   rO   r   r   r   rV   citation_datas          r   _parse_citation_startz3CohereV2ModelResponseIterator._parse_citation_start  s    yy$"%))Ir*KKR0	"w2 }}UA.!fb1$==B7!fn=M  -11r   c           	      n   |j                  di       }|j                  di       }d}|j                  dd      }d}|j                  di       }|rc|j                  di       }t        |j                  d	d
      |j                  dd
      |j                  d	d
      |j                  dd
      z         }|||fS )z:Parse message-end events to extract finish info and usage.r   r   TrU   stopNrY   tokensinput_tokensr   output_tokens)prompt_tokenscompletion_tokenstotal_tokens)r\   r   )	r   rO   r   r   rT   rU   rY   
usage_datatokens_datas	            r   _parse_message_endz0CohereV2ModelResponseIterator._parse_message_end  s    yy$"%		/6:YYw+
$..26K,)oona@"-///1"E(__^Q?+//RacdBeeE M500r   c           	         	 d}d}d}d}d}d}t        |j                  dd            }|j                  dd      }	|j                  dd      }
|	dk(  r| j                  |      }n_|	d	k(  r| j                  |      }nH|
d
k(  r| j	                  |      }n1|
dk(  r| j                  |      }n|
dk(  r| j                  |      \  }}}d|v r|i }|d   |d<   t        |||||||      S # t        $ r}t        d| d|       d}~ww xY w)a  
        Parse Cohere v2 streaming chunks.
        
        v2 format:
        - Content: chunk.type == "content-delta" -> chunk.delta.message.content.text
        - Tool calls: chunk.type == "tool-call-delta" -> chunk.delta.tool_calls
        - Tool plan: chunk.event == "tool-plan-delta" -> chunk.data.delta.message.tool_plan
        - Citations: chunk.event == "citation-start" -> chunk.data.delta.message.citations
        - Finish: chunk.event == "message-end" -> chunk.data.delta.finish_reason
        rQ   NFrR   r   r   eventzcontent-deltaztool-call-deltaztool-plan-deltazcitation-startzmessage-endrV   rW   zFailed to parse v2 chunk: z	, chunk: )
r[   r\   r   r   r   r   r   r   	Exceptionr_   )r   rO   rS   rX   rT   rU   rY   rZ   rR   
chunk_type
event_typero   s               r   ra   z*CohereV2ModelResponseIterator.chunk_parser/  sW   )	OD>BHKM8<E'+$		'1-.E62.J7B/J _,0070066u=00+/+F+Fu+M(//+/+E+Ee+L(},484K4KE4R1]E e#+3/1,8=k8J(5(!'+)A   	O9!IeWMNN	Os   C"C% %	D.C??Dc                     | S r)   r"   rc   s    r   rd   z&CohereV2ModelResponseIterator.__iter__f  re   r   c                    	 | j                   j                         }	 | j                  |      S # t        $ r t        t        $ r}t	        d|       d }~ww xY w# t        $ r t        t        $ r}t	        d| d|       d }~ww xY wrg   ri   rn   s      r   rj   z&CohereV2ModelResponseIterator.__next__i  rp   rq   c                     |}t        |t              r,|j                  d      }|j                  d      }|dk7  r||d }t	        j
                  |      }| j                  |      S )z
        Convert a string chunk to a GenericStreamingChunk for v2

        Note: This is used for Cohere v2 pass through streaming logging
        rs   rt   rG   Nrh   ru   r{   s        r   rm   z@CohereV2ModelResponseIterator.convert_str_chunk_to_generic_chunkx  r~   r   c                 D    | j                   j                         | _        | S r)   r   rc   s    r   r   z'CohereV2ModelResponseIterator.__aiter__  r   r   c                 4  K   	 | j                   j                          d {   }	 | j                  |      S 7 # t        $ r t        t        $ r}t	        d|       d }~ww xY w# t        $ r t        t        $ r}t	        d| d|       d }~ww xY wwrg   r   rn   s      r   r   z'CohereV2ModelResponseIterator.__anext__  r   r   Nr   )r   r   r   __doc__r   r   r   r<   r:   r   r
   r   r   r   r   r   r   r   ra   rd   rj   rm   r   r   r"   r   r   r   r      s    < RW#/3#@H#	$ 	3 	D X>Y5Z D Xd^ 4 HTN "1 1tS(KcBd7d1e 1&4O$ 4O+@ 4OnV2 2@U 2"Vr   r   r)   )r]   typingr   r   r   r    litellm.llms.base_llm.base_utilsr   )litellm.llms.base_llm.chat.transformationr   litellm.types.llms.openair	   litellm.types.utilsr
   r   r   r   r   r   r<   r:   r2   rB   r   r"   r   r   <module>r      s     1 1 = F 6 C" C
B& BR " #$ 	
 c] 
>bV bVHV Vr   