CoolFace
Apppublic

chuckc/Confluence_Sum

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
app.py439 linesDownload Raw Back to root
1import streamlit as st
2from atlassian import Confluence
3import os
4from dotenv import load_dotenv
5from bs4 import BeautifulSoup
6import datetime
7from langchain.schema import Document
8from langchain.chains.summarize import load_summarize_chain
9from langchain_community.llms import Ollama
10from pptx import Presentation
11from pptx.util import Inches, Pt
12from pptx.enum.text import PP_ALIGN
13from langchain.text_splitter import RecursiveCharacterTextSplitter
14from openai import OpenAI
15import time
16import re
17from transformers import GPT2Tokenizer
18
19# Load environment variables
20load_dotenv()
21
22def main():
23    st.title("Confluence Page Summarizer")
24
25    # Add this function to handle API key inputs
26    def get_api_keys():
27        with st.sidebar:
28            st.header("API Configuration")
29            nvidia_conf_user_id = st.text_input("NVIDIA Confluence User ID", value=os.getenv("NVIDIA_CONF_USER_ID", ""))
30            nvidia_conf_api_key = st.text_input("NVIDIA Confluence API Key", type="password", value=os.getenv("NVIDIA_CONF_API_KEY", ""))
31            nvidia_api_key = st.text_input("NVIDIA AI API Key", type="password", value=os.getenv("NVIDIA_API_KEY", ""))
32            #nvidia_conf_user_id = st.text_input("NVIDIA Confluence User ID")
33            #nvidia_conf_api_key = st.text_input("NVIDIA Confluence API Key", type="password")
34            #nvidia_api_key = st.text_input("NVIDIA AI API Key", type="password")           
35            
36            st.markdown("---")
37            st.subheader("Need to register?")
38            st.markdown("""
39            For NVIDIA AI Foundation Models:
40            1. Visit [NVIDIA AI Playground](https://build.nvidia.com/explore/discover)
41            2. Sign up or log in
42            3. Get your API key from the dashboard
43            """)
44
45            st.markdown("---")
46            st.markdown("""
47            For NVIDIA Confluence API Key:
48            1. Visit [Confluence API Guide](https://confluence.nvidia.com/pages/viewpage.action?spaceKey=VLSIPR&title=How+to+use+Confluence+REST+API)
49            2. Follow the instructions to generate your API key
50            """)
51            
52        return nvidia_conf_user_id, nvidia_conf_api_key, nvidia_api_key
53
54    # Get API keys from sidebar
55    nvidia_conf_user_id, nvidia_conf_api_key, nvidia_api_key = get_api_keys()
56
57    # Initialize Confluence client with the new inputs
58    confluence_url = 'https://confluence.nvidia.com/'
59    
60    try:
61        from atlassian import Confluence
62        confluence = Confluence(
63            url=confluence_url,
64            username=nvidia_conf_user_id,
65            password=nvidia_conf_api_key
66        )
67
68    except ImportError:
69        st.error("Failed to import Atlassian Confluence. Please check your installation.")
70        return
71
72    try:
73        from openai import OpenAI
74        # Initialize OpenAI client for NVIDIA AI Foundation Models
75        client = OpenAI(
76            base_url="https://integrate.api.nvidia.com/v1",
77            api_key=nvidia_api_key
78        )
79    except ImportError:
80        st.error("Failed to import OpenAI. Please check your installation.")
81        return
82
83    def get_content_by_url(url):
84        try:
85            if 'viewpage.action' in url:
86                params = dict(param.split('=') for param in url.split('?')[1].split('&'))
87                if 'pageId' in params:
88                    page_id = params['pageId']
89                    page = confluence.get_page_by_id(page_id, expand='body.storage')
90                elif 'spaceKey' in params and 'title' in params:
91                    space = params['spaceKey']
92                    title = params['title']
93                    page = confluence.get_page_by_title(space, title, expand='body.storage')
94                else:
95                    return "Invalid URL format", None
96            else:
97                # Existing logic for other URL formats
98                parts = url.split('/')
99                space = parts[-2]
100                title = parts[-1].replace('+', ' ')  # Replace '+' with space
101                page = confluence.get_page_by_title(space, title, expand='body.storage')
102            
103            if isinstance(page, dict) and 'body' in page and 'storage' in page['body']:
104                html_content = page['body']['storage']['value']
105                formatted_content = format_html_content(html_content)
106                return formatted_content, page.get('title', 'Unknown Title')
107            else:
108                st.error(f"Unexpected response from Confluence API: {page}")
109                return "Failed to retrieve page content", None
110        except Exception as e:
111            st.error(f"Error retrieving page content: {str(e)}")
112            return f"Error: {str(e)}", None
113
114    def format_html_content(html_content):
115        soup = BeautifulSoup(html_content, 'html.parser')
116        
117        # Function to extract text from an element
118        def extract_text(element):
119            if element.name in ['h1', 'h2', 'h3', 'h4', 'h5', 'h6']:
120                return f"\n{element.text.strip()}\n"
121            elif element.name == 'a' and element.get('href'):
122                return f"{element.text.strip()} ({element.get('href')})"
123            elif element.name in ['ul', 'ol']:
124                return '\n' + '\n'.join(f"- {li.text.strip()}" for li in element.find_all('li')) + '\n'
125            elif element.name == 'pre':
126                return f"\n```\n{element.text.strip()}\n```\n"
127            elif element.name == 'table':
128                rows = []
129                for row in element.find_all('tr'):
130                    cells = [cell.text.strip() for cell in row.find_all(['th', 'td'])]
131                    rows.append(' | '.join(cells))
132                return '\n' + '\n'.join(rows) + '\n'
133            else:
134                return element.text.strip() + ' '
135
136        # Extract all text content in order
137        content = []
138        root = soup.body if soup.body else soup  # Use soup if body is None
139        for element in root.descendants:
140            if isinstance(element, str) and element.strip():
141                content.append(element.strip() + ' ')
142            elif element.name:
143                content.append(extract_text(element))
144
145        return ''.join(content)
146
147    def convert_to_langchain_document(content, url, title):
148        metadata = {
149            "source": url,
150            "title": title
151        }
152        return Document(
153            page_content=content,
154            metadata=metadata
155        )
156
157    def count_tokens(text):
158        tokenizer = GPT2Tokenizer.from_pretrained("gpt2")
159        return len(tokenizer.encode(text))
160
161    def summarize_document(document):
162        # Create a text splitter
163        text_splitter = RecursiveCharacterTextSplitter(
164            chunk_size=100000,
165            chunk_overlap=2000,
166            length_function=len
167        )
168        
169        # Split the document into chunks
170        chunks = text_splitter.split_text(document.page_content)
171        total_chunks = len(chunks)
172        
173        print(f"Document split into {total_chunks} chunks.")
174        
175        total_input_tokens = 0
176        total_output_tokens = 0
177        
178        if total_chunks == 1:
179            print("Document is small enough for single summarization.")
180            prompt = f"Provide a concise summary of the following text:\n\n{document.page_content}\n\nSummary:"
181            
182            input_tokens = count_tokens(prompt)
183            total_input_tokens += input_tokens
184            
185            start_time = time.time()
186            completion = client.chat.completions.create(
187                model="meta/llama-3.1-405b-instruct",
188                messages=[{"role": "user", "content": prompt}],
189                temperature=0.2,
190                top_p=0.7,
191                max_tokens=20000
192            )
193            end_time = time.time()
194            
195            summary = completion.choices[0].message.content
196            output_tokens = count_tokens(summary)
197            total_output_tokens += output_tokens
198            
199            print(f"Summary generated in {end_time - start_time:.2f} seconds")
200            print(f"Input tokens: {input_tokens}, Output tokens: {output_tokens}")
201            return summary, total_input_tokens, total_output_tokens
202        
203        else:
204            print("Starting chunked summarization...")
205            summaries = []
206            for i, chunk in enumerate(chunks, 1):
207                print(f"Processing chunk {i}/{total_chunks}...")
208                start_time = time.time()
209                
210                prompt = f"Summarize the following text:\n\n{chunk}\n\nSummary:"
211                
212                input_tokens = count_tokens(prompt)
213                total_input_tokens += input_tokens
214                
215                completion = client.chat.completions.create(
216                    model="meta/llama-3.1-405b-instruct",
217                    messages=[{"role": "user", "content": prompt}],
218                    temperature=0.2,
219                    top_p=0.7,
220                    max_tokens=20000
221                )
222                
223                chunk_summary = completion.choices[0].message.content
224                summaries.append(chunk_summary)
225                
226                output_tokens = count_tokens(chunk_summary)
227                total_output_tokens += output_tokens
228                
229                end_time = time.time()
230                print(f"Chunk {i} summarized in {end_time - start_time:.2f} seconds")
231                print(f"Input tokens: {input_tokens}, Output tokens: {output_tokens}")
232            
233            print("All chunks summarized. Generating final summary...")
234            
235            # Combine summaries
236            combined_summary = " ".join(summaries)
237            
238            # Final summarization of combined summaries
239            final_prompt = f"Provide a concise summary of the following text:\n\n{combined_summary}\n\nFinal Summary:"
240            final_input_tokens = count_tokens(final_prompt)
241            total_input_tokens += final_input_tokens
242            
243            start_time = time.time()
244            final_completion = client.chat.completions.create(
245                model="meta/llama-3.1-405b-instruct",
246                messages=[{"role": "user", "content": final_prompt}],
247                temperature=0.2,
248                top_p=0.7,
249                max_tokens=10000
250            )
251            end_time = time.time()
252            
253            final_summary = final_completion.choices[0].message.content
254            final_output_tokens = count_tokens(final_summary)
255            total_output_tokens += final_output_tokens
256            
257            print(f"Final summary generated in {end_time - start_time:.2f} seconds")
258            print(f"Final input tokens: {final_input_tokens}, Final output tokens: {final_output_tokens}")
259            print(f"Total input tokens: {total_input_tokens}, Total output tokens: {total_output_tokens}")
260            
261            return final_summary, total_input_tokens, total_output_tokens
262
263    def create_summary_slide(title, summary, url):
264        prs = Presentation()
265        slide_layout = prs.slide_layouts[1]  # Using the bullet slide layout
266        
267        def add_slide(slide_number):
268            slide = prs.slides.add_slide(slide_layout)
269            
270            # Set the slide title
271            title_shape = slide.shapes.title
272            title_shape.text = f"{title}"
273            title_shape.text_frame.paragraphs[0].font.size = Pt(24)
274            
275            # Add slide number
276            slide_number_box = slide.shapes.add_textbox(Inches(9), Inches(0.1), Inches(1), Inches(0.5))
277            slide_number_box.text_frame.text = f"Slide {slide_number}"
278            slide_number_box.text_frame.paragraphs[0].alignment = PP_ALIGN.RIGHT
279            slide_number_box.text_frame.paragraphs[0].font.size = Pt(14)
280            
281            # Add the summary placeholder
282            body_shape = slide.shapes.placeholders[1]
283            tf = body_shape.text_frame
284            tf.text = "Summary:"
285            tf.paragraphs[0].font.size = Pt(18)
286            tf.paragraphs[0].font.bold = True
287            
288            return tf
289        
290        # Split the summary into sentences
291        sentences = summary.split('. ')
292        
293        slide_number = 1
294        tf = add_slide(slide_number)
295        
296        for index, sentence in enumerate(sentences, 1):
297            p = tf.add_paragraph()
298            p.text = f"{index}. {sentence.strip()}"
299            p.font.size = Pt(14)
300            p.level = 0
301            
302            # If the text frame is full, start a new slide
303            if len(tf.paragraphs) > 8:
304                slide_number += 1
305                tf = add_slide(slide_number)
306        
307        # Add the source URL to the last slide
308        left = Inches(0.5)
309        top = Inches(6.8)
310        width = Inches(9)
311        height = Inches(0.3)
312        textbox = prs.slides[-1].shapes.add_textbox(left, top, width, height)
313        textbox.text_frame.text = f"Source: {url}"
314        textbox.text_frame.paragraphs[0].alignment = PP_ALIGN.CENTER
315        textbox.text_frame.paragraphs[0].font.size = Pt(10)
316        textbox.text_frame.paragraphs[0].font.italic = True
317        
318        return prs
319
320    # Initialize session state
321    if 'summary' not in st.session_state:
322        st.session_state.summary = None
323    if 'files' not in st.session_state:
324        st.session_state.files = {}
325    if 'last_url' not in st.session_state:
326        st.session_state.last_url = ""
327    if 'input_tokens' not in st.session_state:
328        st.session_state.input_tokens = 0
329    if 'output_tokens' not in st.session_state:
330        st.session_state.output_tokens = 0
331    
332    # Input for Confluence URL
333    url = st.text_input(
334        "Enter Confluence page URL:",
335        placeholder="https://confluence.nvidia.com/display/SPACE/Page+Title"
336    )
337
338    # Clear summary if a new URL is entered
339    if url != st.session_state.last_url:
340        st.session_state.summary = None
341        st.session_state.files = {}
342        st.session_state.last_url = url
343        st.session_state.input_tokens = 0
344        st.session_state.output_tokens = 0
345
346    if st.button("Summarize"):
347        if url:
348            with st.spinner("Retrieving and summarizing content..."):
349                content, title = get_content_by_url(url)
350                
351                if title and not content.startswith("Error:"):
352                    document = convert_to_langchain_document(content, url, title)
353                    summary, input_tokens, output_tokens = summarize_document(document)
354
355                    st.session_state.summary = summary
356                    st.session_state.input_tokens = input_tokens
357                    st.session_state.output_tokens = output_tokens
358
359                    # Create and save files
360                    folder_name = "confluence_content"
361                    os.makedirs(folder_name, exist_ok=True)
362
363                    timestamp = datetime.datetime.now().strftime("%Y%m%d_%H%M%S")
364                    
365                    # Sanitize the filename
366                    sanitized_title = re.sub(r'[^\w\-_\. ]', '_', title)
367                    filename_base = f"{sanitized_title}_{timestamp}"
368
369                    # Save PowerPoint
370                    prs = create_summary_slide(title, summary, url)
371                    ppt_filename = f"{filename_base}.pptx"
372                    ppt_output_file = os.path.join(folder_name, ppt_filename)
373                    prs.save(ppt_output_file)
374                    
375                    # Save content
376                    content_filename = f"{filename_base}.txt"
377                    content_output_file = os.path.join(folder_name, content_filename)
378                    with open(content_output_file, 'w', encoding='utf-8') as file:
379                        file.write(f"URL: {url}\n\n{content}")
380                    
381                    # Save summary
382                    summary_filename = f"summary_{filename_base}.txt"
383                    summary_output_file = os.path.join(folder_name, summary_filename)
384                    with open(summary_output_file, 'w', encoding='utf-8') as file:
385                        file.write(f"Summary of {url}\n\n{summary}")
386
387                    # Store file paths in session state
388                    st.session_state.files = {
389                        'ppt': (ppt_output_file, ppt_filename),
390                        'content': (content_output_file, content_filename),
391                        'summary': (summary_output_file, summary_filename)
392                    }
393                else:
394                    st.error(f"Failed to retrieve the page: {content}")
395        else:
396            st.warning("Please enter a Confluence page URL.")
397
398    # Display summary, token counts, and download buttons if available
399    if st.session_state.summary is not None:
400        st.subheader("Summary")
401        st.write(st.session_state.summary)
402
403        # Display token counts
404        st.subheader("Token Usage")
405        col1, col2, col3 = st.columns(3)
406        with col1:
407            st.metric("Input Tokens", st.session_state.input_tokens)
408        with col2:
409            st.metric("Output Tokens", st.session_state.output_tokens)
410        with col3:
411            st.metric("Total Tokens", st.session_state.input_tokens + st.session_state.output_tokens)
412
413        # Provide download buttons with icons
414        if st.session_state.files:
415            st.subheader("Download Files")
416            for file_type, (file_path, file_name) in st.session_state.files.items():
417                with open(file_path, "rb") as file:
418                    if file_type == 'ppt':
419                        mime = "application/vnd.openxmlformats-officedocument.presentationml.presentation"
420                        label = "Download PowerPoint Summary"
421                        icon = "๐Ÿ“Š"  # Chart icon for PowerPoint
422                    elif file_type == 'content':
423                        mime = "text/plain"
424                        label = "Download Full Content"
425                        icon = "๐Ÿ“„"  # Page icon for full content
426                    else:  # summary
427                        mime = "text/plain"
428                        label = "Download Summary Text"
429                        icon = "๐Ÿ“"  # Memo icon for summary
430                    
431                    st.download_button(
432                        label=f"{icon} {label}",
433                        data=file,
434                        file_name=file_name,
435                        mime=mime
436                    )
437
438if __name__ == "__main__":
439    main()