Team Ai
Apppublic

Hariprita/nl2sql-openenv

sourceHugging Faceupdated 6mo agoView on Hugging Face
0likes
environment.cpython-311.pyc103 linesDownload Raw Back to __pycache__
1�

2���i'��l�ddlZddlmZddlmZmZmZddlmZm	Z	m3Z4mZe	e5ed�ZGd�d��Z
dS)	�N)�Optional�)�create_connection�get_schema_ddl�
execute_query)�TASKS�grade_simple_select�grade_join_aggregation�grade_window_ranking)�
simple_select�join_aggregation�window_rankingc�|�eZdZdZd�Zdefd�Zdedefd�Ze	defd���Z6edefd���Zd	edefd7�Z
dS)�SQLAgentEnvironment�c���d|_t��|_d|_d|_d|_t
tj����|_	d|_8d|_dS)Nrr�F)�_connr�_schema_ddl�	_task_idx�_attempt�_step_count�str�uuid�uuid4�_episode_id�_cumulative_reward�_done)�selfs �6C:\Users\Hariprita\sql_agent_env\server\environment.py�__init__zSQLAgentEnvironment.__init__sV����9�)�+�+��������
�����t�z�|�|�,�,���"%�����10�11�12��returnc�l�t��|_d|_d|_d|_tt
j����|_d|_	d|_13td}|j|ddddd|d�d	�|d14|d|j|j
|�dd��d�S)
z:Start a fresh episode. Returns the first observation dict.rrrF�question�z15Task 1/3 (�16difficultyz%): ready. Write your first SQL query.�id�hint��schemar%�result�reward�done�feedback�task_id�task_difficulty�attempt�max_attemptsr))rrrrrrrrrrrrr�MAX_ATTEMPTS_PER_TASK�get)r�tasks  r �resetzSQLAgentEnvironment.resets���&�(�(��17������
�����t�z�|�|�,�,���"%�����18��Q�x��#�/�#�J�/�!�"�$�e�D��,>�e�e�e�#�D�z�#�L�1�#�}�#�9�#�x�x���3�3�19�20�	21r"�	sql_queryc��|jr|�d��St|j}t	|j|��\}}}t|d}||||��\}}|xjdz
c_|xj|z
c_|dkp|j	|j22k}	|	r|xjdz
c_d|_	n|xj	dz
c_	|jtt��k}23|24|_|�|||��}|25s�t|j}|	r|�d|jdz�d|d�d�n|�d	|j	�d26|j27�d�}
|j
|d||d
|
|d|d|j	|j28|�dd��d�S|�dtt���d|jd��}
|j
|d||d|
|d|d|j	|j29dd�S)zGExecute sql_query, grade it, advance episode state, return observation.z>Episode already complete. Call reset() to start a new episode.r(rgffffff�?z -> Moving to task z/3 (r'z).z	 Attempt �/�.r%Fr)r&r*z All z$ tasks complete! Cumulative reward: z.2fT)r�_make_terminal_obsrrrr�GRADERSrrrr4�len�_format_resultrr5)rr8r6�rows�columns�error�graderr-r/�advancer.�30result_str�	next_task�
feedback_fulls              r �stepzSQLAgentEnvironment.step8sg���:�	��*�*�P���
��T�^�$��,�T�Z��C�C���g�u���d��$��!�6�$���7�7�������A�������6�)����S�=�R�d�m�t�7Q�&Q���	��N�N�a��N�N��D�M�M��M�M�Q��M�M��~��U���+����31��(�(��w��>�>�32��%	��d�n�-�I��Y�8�c�c����0B�c�c�	�R^�H_�c�c�c�c� �X�X�4�=�X�X�4�;U�X�X�X�
�$(�#3�#,�Z�#8�#-�#)�#(�#0�#,�T�?�#,�\�#:�#'�=�#'�#=�#,�=�=���#<�#<���
��D�D�#�e�*�*�D�D�&*�&=�C�D�D�
�33$(�#3�#'�34�#3�#-�#)�#'�#0�#'��:�#'��#5�#'�=�#'�#=�#%���
r"c���t|jtt��dz35��}|j|jt|d|jtt��|j|jd�S)Nrr()�36episode_id�37step_count�current_task_id�current_task_idx�total_tasks�cumulative_rewardr.)�minrr>rrrrr)r�safe_idxs  r �statezSQLAgentEnvironment.states]���t�~�s�5�z�z�A�~�6�6�� $� 0� $� 0� %�h��� 5� $�� #�E�38�39�!%�!8� $�40�41�42�	43r"c�>��|rd|��S|sdSd����}dtt|��d��z}d��fd�|dd�D����}t|��dkrd	t|���d44�nd}|�d|�d|�|��S)NzERROR: z(empty result set)� | �-�45�46c3�\��K�|]%�d��fd��D����V��&dS)rTc3�^�K�|]'}t��|d����V��(dS)r&N)rr5)�.0�c�rs  �r �	<genexpr>z?SQLAgentEnvironment._format_result.<locals>.<genexpr>.<genexpr>�s7�����:�:�Q�s�1�5�5��B�<�<�(�(�:�:�:�:�:�:r"N)�join)rZr\rAs @�r r]z5SQLAgentEnvironment._format_result.<locals>.<genexpr>�sY������47�48��
�J�J�:�:�:�:�'�:�:�:�:�:�49�50�51�52�53�54r"�z55... (z rows total)r&)r^�maxr>)r@rArB�header�	separator�rows_str�suffixs `     r r?z"SQLAgentEnvironment._format_result�s�����	%�$�U�$�$�$��	(�'�'����G�$�$���#�c�&�k�k�2�.�.�.�	��9�9�56�57�58�59��"�1�"�X�60�61�62�63�64��7:�$�i�i�!�m�m�2�3�t�9�9�2�2�2�2����;�;�I�;�;��;�6�;�;�;r"�messagec���t|jtt��dz65��}t|}|j|d|dd||d|d|j|jdd�S)	Nrr%rTr(r'r&r*)rPrr>rrrr4)rrerQr6s    r r<z&SQLAgentEnvironment._make_terminal_obs�sl���t�~�s�5�z�z�A�~�6�6���X���#�/�#�J�/�&�"�#�&�#�D�z�#�L�1�#�}�#�9�!�66�67�	68r"N)�__name__�69__module__�__qualname__r4r!�dictr7rrH�propertyrR�staticmethodr?r<�r"r rrs������������70�t�71�72�73�74�2E�c�E�d�E�E�E�E�N�7576�t�7778�7980�8182��X�8384� �<��<�<�<��\�<�85�#�86�$�87�88�89�90�91�92r"r)r�typingr�databaserrr�tasksrr	r93rr=rrmr"r �<module>rqs�������������F�F�F�F�F�F�F�F�F�F�[�[�[�[�[�[�[�[�[�[�[�[�+�.�+����`94�`95�`96�`97�`98�`99�`100�`101�`102�`103r"